| 142 | |
| 143 | |
| 144 | def _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 | |