Changed max_edge in example/plot_rag.py

This commit is contained in:
Vighnesh Birodkar
2014-06-28 11:06:30 +05:30
parent f71ad51f57
commit dd67b3fce7
3 changed files with 11 additions and 19 deletions
+6 -18
View File
@@ -9,32 +9,20 @@ weight of all the edges incident on it can be updated by a user defined
function `weight_func`. function `weight_func`.
The default behaviour is to use the smaller edge weight incase of a conflict. The default behaviour is to use the smaller edge weight incase of a conflict.
THe example below also shows how to use a custom function to take the larger The example below also shows how to use a custom function to take the larger
weight instead. weight instead.
""" """
from skimage.graph import rag from skimage.graph import rag
import networkx as nx import networkx as nx
from matplotlib import pyplot as plt from matplotlib import pyplot as plt
import numpy as np
def max_edge(g, src, dst, neighbor): def max_edge(g, src, dst, n):
try: w1 = g[n].get(src, {'weight': -np.inf})['weight']
w1 = g.edge[src][neighbor]['weight'] w2 = g[n].get(dst, {'weight': -np.inf})['weight']
except KeyError: return max(w1, w2)
w1 = None
try:
w2 = g.edge[dst][neighbor]['weight']
except KeyError:
w2 = None
if w1 is None:
return w2
elif w2 is None:
return w1
else:
return max(w1, w2)
def display(g, title): def display(g, title):
+1 -1
View File
@@ -17,7 +17,7 @@ def min_weight(g, src, dst, n):
src, dst : int src, dst : int
The verices in `g` to be merged. The verices in `g` to be merged.
n : int n : int
A neighbor of `src` or `dst` or both A neighbor of `src` or `dst` or both.
Returns Returns
------- -------
+4
View File
@@ -23,10 +23,14 @@ def test_rag_merge():
gc = g.copy() gc = g.copy()
# We merge nodes and ensure that the minimum weight is chosen
# when there is a conflict.
g.merge_nodes(0, 2) g.merge_nodes(0, 2)
assert g.edge[1][2]['weight'] == 10 assert g.edge[1][2]['weight'] == 10
assert g.edge[2][3]['weight'] == 30 assert g.edge[2][3]['weight'] == 30
# We specify `max_edge` as `weight_func` as ensure that maximum
# weight is chosen in case on conflict
gc.merge_nodes(0, 2, weight_func=max_edge) gc.merge_nodes(0, 2, weight_func=max_edge)
assert gc.edge[1][2]['weight'] == 20 assert gc.edge[1][2]['weight'] == 20
assert gc.edge[2][3]['weight'] == 40 assert gc.edge[2][3]['weight'] == 40