MCPcopy Create free account
hub / github.com/M-Nauta/ProtoTree / _gen_dot_edges

Function _gen_dot_edges

util/visualize.py:144–164  ·  view source on GitHub ↗
(node: Node, classes:tuple)

Source from the content-addressed store, hash-verified

142
143
144def _gen_dot_edges(node: Node, classes:tuple):
145 if isinstance(node, Branch):
146 edge_l, targets_l = _gen_dot_edges(node.l, classes)
147 edge_r, targets_r = _gen_dot_edges(node.r, classes)
148 str_targets_l = ','.join(str(t) for t in targets_l) if len(targets_l) > 0 else ""
149 str_targets_r = ','.join(str(t) for t in targets_r) if len(targets_r) > 0 else ""
150 s = '{} -> {} [label="Absent" fontsize=10 tailport="s" headport="n" fontname=Helvetica];\n {} -> {} [label="Present" fontsize=10 tailport="s" headport="n" fontname=Helvetica];\n'.format(node.index, node.l.index,
151 node.index, node.r.index)
152 return s + edge_l + edge_r, sorted(list(set(targets_l + targets_r)))
153 if isinstance(node, Leaf):
154 if node._log_probabilities:
155 ws = copy.deepcopy(torch.exp(node.distribution()).cpu().detach().numpy())
156 else:
157 ws = copy.deepcopy(node.distribution().cpu().detach().numpy())
158 argmax = np.argmax(ws)
159 targets = [argmax] if argmax.shape == () else argmax.tolist()
160 class_targets = copy.deepcopy(targets)
161 for i in range(len(targets)):
162 t = targets[i]
163 class_targets[i] = classes[t]
164 return '', class_targets
165

Callers 1

gen_visFunction · 0.85

Calls 1

distributionMethod · 0.80

Tested by

no test coverage detected