diff --git a/backends/vulkan/serialization/vulkan_graph_serialize.py b/backends/vulkan/serialization/vulkan_graph_serialize.py index 665c23b8df2..8707ad53f50 100644 --- a/backends/vulkan/serialization/vulkan_graph_serialize.py +++ b/backends/vulkan/serialization/vulkan_graph_serialize.py @@ -110,7 +110,7 @@ def convert_to_flatbuffer(vk_graph: VkGraph) -> bytes: json_path = os.path.join(d, "schema.json") with open(json_path, "wb") as json_file: json_file.write(vk_graph_json.encode("ascii")) - _flatc_compile(d, schema_path, json_path) + _flatc_compile(d, schema_path, json_path, ["--force-defaults"]) output_path = os.path.join(d, "schema.bin") with open(output_path, "rb") as output_file: return output_file.read() diff --git a/backends/vulkan/test/test_serialization.py b/backends/vulkan/test/test_serialization.py index 540b86ace82..8b95b1df2a1 100644 --- a/backends/vulkan/test/test_serialization.py +++ b/backends/vulkan/test/test_serialization.py @@ -372,6 +372,14 @@ def test_serialize_deserialize_non_finite_scalars(self) -> None: ] ) + def test_serialize_deserialize_signed_zero(self) -> None: + values = [ + VkValue(Double(-0.0)), + VkValue(Double(0.0)), + VkValue(DoubleList([-0.0, 0.0])), + ] + self.assertEqual(repr(self._round_trip(values).values), repr(values)) + def test_serialize_deserialize_non_finite_floats_in_list(self) -> None: # json only emits a float as a chunk of its own inside an object; in a # list the chunk carries the delimiter with it, so a rewrite that works diff --git a/exir/_serialize/_flatbuffer.py b/exir/_serialize/_flatbuffer.py index eaf7f296e55..4d015b84a80 100644 --- a/exir/_serialize/_flatbuffer.py +++ b/exir/_serialize/_flatbuffer.py @@ -369,7 +369,12 @@ def convert(value: object) -> object: return json.dumps(convert(json.loads(content))).encode("utf-8") -def _flatc_compile(output_dir: str, schema_path: str, json_path: str) -> None: +def _flatc_compile( + output_dir: str, + schema_path: str, + json_path: str, + flatc_additional_args: Optional[List[str]] = None, +) -> None: """Serializes JSON data to a binary flatbuffer file. Args: @@ -381,6 +386,7 @@ def _flatc_compile(output_dir: str, schema_path: str, json_path: str) -> None: matches the schema. Rewritten in place if it contains non-finite floats; flatc derives the output filename from this path, so the data cannot be sanitized into a differently-named file. + flatc_additional_args: Additional options passed to flatc. """ with open(json_path, "rb") as json_file: content = json_file.read() @@ -392,6 +398,7 @@ def _flatc_compile(output_dir: str, schema_path: str, json_path: str) -> None: _run_flatc( [ "--binary", + *(flatc_additional_args or []), "-o", output_dir, schema_path,