( # noqa: C901
title: str,
path: str,
edge_ep: exir.ExirExportedProgram,
numeric_results: pandas.core.frame.DataFrame,
comparator: QcomNumericalComparatorBase,
)
| 74 | |
| 75 | |
| 76 | def export_svg( # noqa: C901 |
| 77 | title: str, |
| 78 | path: str, |
| 79 | edge_ep: exir.ExirExportedProgram, |
| 80 | numeric_results: pandas.core.frame.DataFrame, |
| 81 | comparator: QcomNumericalComparatorBase, |
| 82 | ): |
| 83 | def get_node_style(is_valid_score: bool): |
| 84 | template = { |
| 85 | "shape": "record", |
| 86 | "style": '"filled,rounded"', |
| 87 | "fontcolor": "#000000", |
| 88 | } |
| 89 | |
| 90 | if is_valid_score is None: |
| 91 | template["fillcolor"] = "LemonChiffon1" # No match between QNN and CPU |
| 92 | elif is_valid_score: |
| 93 | template["fillcolor"] = "DarkOliveGreen3" # Good accuracy |
| 94 | else: |
| 95 | template["fillcolor"] = "Coral1" # Bad accuracy |
| 96 | |
| 97 | return template |
| 98 | |
| 99 | pydot_graph = pydot.Dot(graph_type="graph") |
| 100 | node_map = {} |
| 101 | |
| 102 | # Create node |
| 103 | for node in edge_ep.graph_module.graph.nodes: |
| 104 | # These are just nodes before fold_quant and still there |
| 105 | if len(node.users) == 0 and node.op == "placeholder": |
| 106 | continue |
| 107 | |
| 108 | pytorch_layout = get_pytorch_layout_info(node) |
| 109 | scale_zero_point = get_scale_zero_point(node) |
| 110 | scale = scale_zero_point["scale(s)"] |
| 111 | zero_point = scale_zero_point["zero_point(s)"] |
| 112 | |
| 113 | node_label = "{" |
| 114 | node_label += f"name=%{node.name}" + r"\n" |
| 115 | node_label += f"|op_code={node.op}" + r"\n" |
| 116 | node_label += f"|target={typename(node.target)}" + r"\n" |
| 117 | node_label += f"|num_users={len(node.users)}" + r"\n" |
| 118 | node_label += f"|pytorch_layout={pytorch_layout}" + r"\n" |
| 119 | node_label += f"|scale(s)={scale}" + r"\n" |
| 120 | node_label += f"|zero_point(s)={zero_point}" + r"\n" |
| 121 | |
| 122 | is_valid_score = None |
| 123 | if debug_handle := node.meta.get(DEBUG_HANDLE_KEY, None): |
| 124 | node_label += f"|debug_handle={debug_handle}" + r"\n" |
| 125 | debug_handle = (debug_handle,) |
| 126 | if debug_handle in numeric_results.index: |
| 127 | score = numeric_results.loc[[debug_handle], "gap"].iat[0][0] |
| 128 | assert isinstance( |
| 129 | score, float |
| 130 | ), f"Expecting QcomNumericalComparatorBase element_compare to return float, but get {type(score)}." |
| 131 | node_label += f"|{comparator.metric_name()}={score:.3f}" + r"\n" |
| 132 | is_valid_score = comparator.is_valid_score(score) |
| 133 | node_label += f"|is_valid_score={is_valid_score}" + r"\n" |
no test coverage detected