MCPcopy Create free account
hub / github.com/pytorch/executorch / export_svg

Function export_svg

backends/qualcomm/debugger/format_outputs.py:76–170  ·  view source on GitHub ↗
(  # noqa: C901
    title: str,
    path: str,
    edge_ep: exir.ExirExportedProgram,
    numeric_results: pandas.core.frame.DataFrame,
    comparator: QcomNumericalComparatorBase,
)

Source from the content-addressed store, hash-verified

74
75
76def 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"

Callers 1

generate_resultsMethod · 0.85

Calls 11

get_pytorch_layout_infoFunction · 0.85
get_scale_zero_pointFunction · 0.85
typenameFunction · 0.85
get_node_styleFunction · 0.85
add_nodeMethod · 0.80
keysMethod · 0.80
infoMethod · 0.80
getMethod · 0.45
metric_nameMethod · 0.45
is_valid_scoreMethod · 0.45
runMethod · 0.45

Tested by

no test coverage detected