diff --git a/devtools/etdump/serialize.py b/devtools/etdump/serialize.py index 00fd1e70689..fc656057e28 100644 --- a/devtools/etdump/serialize.py +++ b/devtools/etdump/serialize.py @@ -10,6 +10,7 @@ import importlib.resources as _resources import json import os +import re import tempfile import executorch.devtools.etdump as etdump_package @@ -21,6 +22,10 @@ ETDUMP_FLATCC_SCHEMA_NAME = "etdump_schema_flatcc" SCALAR_TYPE_SCHEMA_NAME = "scalar_type" +# flatc writes non-finite floats as bare inf, -inf, nan and -nan, which +# json.loads rejects. String literals are matched first so they pass through. +_FLATC_NON_FINITE_RE = re.compile(rb'("(?:[^"\\]|\\.)*")|-?\bnan\b|\binf\b') + def _write_schema(d: str, schema_name: str) -> None: schema_path = os.path.join(d, "{}.fbs".format(schema_name)) @@ -39,7 +44,12 @@ def _serialize_from_etdump_to_json(etdump: ETDumpFlatCC) -> str: # from json to etdump def _deserialize_from_json_to_etdump_flatcc(etdump_json: bytes) -> ETDumpFlatCC: - etdump_json = json.loads(etdump_json) + etdump_json = json.loads( + _FLATC_NON_FINITE_RE.sub( + lambda m: m.group(1) or (b"Infinity" if m.group(0) == b"inf" else b"NaN"), + etdump_json, + ) + ) return _json_to_dataclass(etdump_json, ETDumpFlatCC) diff --git a/devtools/etdump/tests/serialize_test.py b/devtools/etdump/tests/serialize_test.py index 5cab3e5b2ba..a2b3dcdaf33 100644 --- a/devtools/etdump/tests/serialize_test.py +++ b/devtools/etdump/tests/serialize_test.py @@ -8,6 +8,7 @@ import difflib import json +import math import unittest from pprint import pformat from typing import List @@ -140,3 +141,23 @@ def test_serialize(self) -> None: ) ), ) + + def test_serialize_non_finite_floats(self) -> None: + program = get_sample_etdump_flatcc() + debug_entry = program.run_data[0].events[3].debug_event.debug_entry + debug_entry.float_value = flatcc.Float(float("inf")) + debug_entry.double_value = flatcc.Double(float("-inf")) + + deserialized_obj = deserialize_from_etdump_flatcc( + serialize_to_etdump_flatcc(program), size_prefixed=False + ) + self.assertEqual(program, deserialized_obj) + + debug_entry.double_value = flatcc.Double(float("nan")) + deserialized_obj = deserialize_from_etdump_flatcc( + serialize_to_etdump_flatcc(program), size_prefixed=False + ) + deserialized_entry = deserialized_obj.run_data[0].events[3].debug_event + self.assertTrue( + math.isnan(deserialized_entry.debug_entry.double_value.double_val) + )