Skip to content

Commit 9dedab4

Browse files
authored
Fix Inspector compare_results crash on non-tensor outputs (pytorch#23276)
### Summary The Inspector's compare_results crashes as soon as one of the outputs isn't a tensor. calculate_mse / calculate_snr / calculate_cosine_similarity return None for non-tensor outputs, and then the print loop does f"{value:>8.5f}" on it and dies with "TypeError: unsupported format string passed to NoneType.__format__". plot=True breaks the same way (max() and plt.bar on None). The CLI hits a related problem: `inspector_cli --compare_results` calls compare_results even when there's no reference output (no ETRecord with reference outputs), so it fails there too. Changes: - print "N/A" for None values instead of formatting them - plot None values as 0 so the bar chart still renders - in inspector_cli, skip the comparison with a short message when reference_output or run_output is missing - return type of compare_results is now Dict[str, List[Optional[float]]], which is what it actually returns ### Test plan Added test_compare_results_with_non_tensor_output in devtools/inspector/tests/inspector_utils_test.py. It calls compare_results with a tensor output and an int output, with plot on and off. It fails on main with the TypeError above and passes with this change. The rest of inspector_utils_test.py passes (90 tests). ufmt and flake8 are clean with the versions from requirements-lintrunner.txt. cc @Gasoonjia @nil-is-all
1 parent 0f008e1 commit 9dedab4

3 files changed

Lines changed: 41 additions & 11 deletions

File tree

‎devtools/inspector/_inspector_utils.py‎

Lines changed: 11 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -402,8 +402,10 @@ def plot_metric(result: List[float], metric_name: str):
402402
plt.clf()
403403
plt.figure(figsize=(8, 6))
404404

405-
x_axis = np.arange(len(result))
406-
bars = plt.bar(x_axis, result, width=0.5)
405+
# Non-tensor outputs have no metric value (None); plot them as 0.
406+
values = [v if v is not None else 0.0 for v in result]
407+
x_axis = np.arange(len(values))
408+
bars = plt.bar(x_axis, values, width=0.5)
407409
plt.grid(True, which="major", axis="y")
408410
num_ticks = len(x_axis) if len(x_axis) > 5 else 5
409411
interval = 1 if num_ticks < 20 else 5
@@ -422,8 +424,8 @@ def plot_metric(result: List[float], metric_name: str):
422424
va="bottom",
423425
)
424426

425-
max_value = max(result) * 1.25
426-
min_value = min(result) * 1.25
427+
max_value = max(values, default=0.0) * 1.25
428+
min_value = min(values, default=0.0) * 1.25
427429

428430
# Cosine similarity has range [-1, 1], so we set y-axis limits accordingly.
429431
if metric_name == "cosine_similarity":
@@ -514,7 +516,7 @@ def compare_results(
514516
run_output: ProgramOutput,
515517
metrics: Optional[List[str]] = None,
516518
plot: bool = False,
517-
) -> Dict[str, List[float]]:
519+
) -> Dict[str, List[Optional[float]]]:
518520
"""
519521
Compares the results of two runs and returns a dictionary of metric names -> lists of metric values. This list matches
520522
the reference output & run output lists, so essentially we compare each pair of values in those two lists.
@@ -546,7 +548,10 @@ def compare_results(
546548
print(supported_metric)
547549
print("-" * 20)
548550
for index, value in enumerate(result):
549-
print(f"{index:<5}{value:>8.5f}")
551+
if value is None:
552+
print(f"{index:<5}{'N/A':>8}")
553+
else:
554+
print(f"{index:<5}{value:>8.5f}")
550555
print("\n")
551556

552557
return results

‎devtools/inspector/inspector_cli.py‎

Lines changed: 13 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -64,12 +64,20 @@ def main() -> None:
6464
inspector.save_data_to_tsv(args.tsv_path)
6565
if args.compare_results:
6666
for event_block in inspector.event_blocks:
67-
if event_block.name == "Execute":
68-
compare_results(
69-
reference_output=event_block.reference_output,
70-
run_output=event_block.run_output,
71-
plot=True,
67+
if event_block.name != "Execute":
68+
continue
69+
if event_block.reference_output is None or event_block.run_output is None:
70+
print(
71+
"Skipping --compare_results: no reference output for this run. "
72+
"Pass --etrecord_path for an ETRecord that holds reference outputs "
73+
"(e.g. generated from a BundledProgram) and an ETDump with run outputs."
7274
)
75+
continue
76+
compare_results(
77+
reference_output=event_block.reference_output,
78+
run_output=event_block.run_output,
79+
plot=True,
80+
)
7381

7482

7583
if __name__ == "__main__":

‎devtools/inspector/tests/inspector_utils_test.py‎

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,6 +30,7 @@
3030
calculate_mse,
3131
calculate_snr,
3232
calculate_time_scale_factor,
33+
compare_results,
3334
convert_to_float_tensor,
3435
create_debug_handle_to_op_node_mapping,
3536
EDGE_DIALECT_GRAPH_KEY,
@@ -229,6 +230,22 @@ def test_compare_results_uint8(self):
229230
self.assertGreater(calculate_snr([a], [b])[0], 30.0)
230231
self.assertAlmostEqual(calculate_cosine_similarity([a], [b])[0], 1.0)
231232

233+
def test_compare_results_with_non_tensor_output(self):
234+
# Non-tensor outputs (e.g. ints) have no metric value and must not
235+
# break printing or plotting.
236+
import matplotlib
237+
238+
matplotlib.use("Agg")
239+
a = torch.rand(4, 4)
240+
b = a.clone()
241+
b[0, 0] += 1e-2
242+
for plot in (False, True):
243+
results = compare_results([a, 3], [b, 3], plot=plot)
244+
for values in results.values():
245+
self.assertEqual(len(values), 2)
246+
self.assertIsNotNone(values[0])
247+
self.assertIsNone(values[1])
248+
232249
def test_merge_overlapping_debug_handles_basic(self):
233250
big_tensor = torch.rand(100, 100)
234251
intermediate_outputs = {

0 commit comments

Comments
 (0)