From f62f62cd37174aad92e71884fcbbbcf4b8db5db8 Mon Sep 17 00:00:00 2001 From: Maya Bennett Date: Wed, 22 Jul 2026 09:28:11 -0700 Subject: [PATCH] Add Python APIs for PTD tensor dictionaries --- extension/flat_tensor/README.md | 15 ++ extension/flat_tensor/serialize/BUCK | 5 + extension/flat_tensor/serialize/__init__.py | 12 ++ extension/flat_tensor/serialize/serialize.py | 155 ++++++++++++++++++- extension/flat_tensor/test/test_serialize.py | 146 +++++++++++++++++ 5 files changed, 332 insertions(+), 1 deletion(-) diff --git a/extension/flat_tensor/README.md b/extension/flat_tensor/README.md index b1d8ed8a8fc..d46de9e26fc 100644 --- a/extension/flat_tensor/README.md +++ b/extension/flat_tensor/README.md @@ -20,6 +20,21 @@ Major usage is to store data outside of the PTE file for clean program-data sepa [serialize.py](https://github.com/pytorch/executorch/blob/main/extension/flat_tensor/serialize/serialize.py) contains the Python serialization and deserialization APIs. +Use `save_ptd` and `load_ptd` to round-trip dictionaries of CPU tensors: + +```python +import torch + +from executorch.extension.flat_tensor.serialize import load_ptd, save_ptd + +save_ptd("state.ptd", {"weight": torch.ones(2, 2)}) +tensors = load_ptd("state.ptd") +``` + +The APIs support strided, non-quantized tensor dtypes represented by the PTD +schema. Contiguous and channels-last strides are preserved; other strided +layouts are saved in contiguous form. + ### Alignment Considerations **Segment alignment**: Data segments are aligned to this value. This is usually some multiple of 2. Specified in the [FlatTensorConfig](https://github.com/pytorch/executorch/blob/main/extension/flat_tensor/serialize/serialize.py#L96). diff --git a/extension/flat_tensor/serialize/BUCK b/extension/flat_tensor/serialize/BUCK index 04fb1146c33..bc69d0889de 100644 --- a/extension/flat_tensor/serialize/BUCK +++ b/extension/flat_tensor/serialize/BUCK @@ -30,6 +30,7 @@ fbcode_target(_kind = runtime.python_library, fbcode_target(_kind = runtime.python_library, name = "serialize", srcs = [ + "__init__.py", "serialize.py", ], resources = [ @@ -39,6 +40,10 @@ fbcode_target(_kind = runtime.python_library, visibility = ["PUBLIC"], deps = [ ":schema", + "//caffe2:torch", "//executorch/exir/_serialize:lib", + "//executorch/exir:scalar_type", + "//executorch/exir:tensor", + "//executorch/exir:tensor_layout", ], ) diff --git a/extension/flat_tensor/serialize/__init__.py b/extension/flat_tensor/serialize/__init__.py index e69de29bb2d..08e59515438 100644 --- a/extension/flat_tensor/serialize/__init__.py +++ b/extension/flat_tensor/serialize/__init__.py @@ -0,0 +1,12 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +from .serialize import load_ptd, save_ptd + +__all__ = [ + "load_ptd", + "save_ptd", +] diff --git a/extension/flat_tensor/serialize/serialize.py b/extension/flat_tensor/serialize/serialize.py index 94303958caa..1eb580fe60c 100644 --- a/extension/flat_tensor/serialize/serialize.py +++ b/extension/flat_tensor/serialize/serialize.py @@ -13,9 +13,10 @@ import os import tempfile from dataclasses import dataclass -from typing import ClassVar, Dict, List, Literal, Optional +from typing import ClassVar, Dict, List, Literal, Optional, Set, Union import executorch.extension.flat_tensor.serialize as serialize_package +import torch from executorch.exir._serialize._cord import Cord from executorch.exir._serialize._dataclass import _DataclassEncoder, _json_to_dataclass @@ -28,6 +29,9 @@ DataSerializer, ) from executorch.exir._serialize.padding import aligned_size, pad_to, padding_required +from executorch.exir.scalar_type import ScalarType +from executorch.exir.tensor import dim_order_from_stride, stride_from_dim_order +from executorch.exir.tensor_layout import TensorLayout from executorch.extension.flat_tensor.serialize.flat_tensor_schema import ( DataSegment, FlatTensor, @@ -45,6 +49,44 @@ # Current version. Keep in sync with c++ version number in serialize. _FLAT_TENSOR_VERSION: int = 0 +# Keep in sync with scalar_type.fbs, which excludes complex types. +_PTD_TO_TORCH_DTYPE: Dict[ScalarType, torch.dtype] = { + ScalarType.BYTE: torch.uint8, + ScalarType.CHAR: torch.int8, + ScalarType.SHORT: torch.int16, + ScalarType.INT: torch.int32, + ScalarType.LONG: torch.int64, + ScalarType.HALF: torch.float16, + ScalarType.FLOAT: torch.float32, + ScalarType.DOUBLE: torch.float64, + ScalarType.BOOL: torch.bool, + ScalarType.QINT8: torch.qint8, + ScalarType.QUINT8: torch.quint8, + ScalarType.QINT32: torch.qint32, + ScalarType.BFLOAT16: torch.bfloat16, + ScalarType.QUINT4x2: torch.quint4x2, + ScalarType.QUINT2x4: torch.quint2x4, + ScalarType.BITS16: torch.bits16, + ScalarType.FLOAT8E5M2: torch.float8_e5m2, + ScalarType.FLOAT8E4M3FN: torch.float8_e4m3fn, + ScalarType.FLOAT8E5M2FNUZ: torch.float8_e5m2fnuz, + ScalarType.FLOAT8E4M3FNUZ: torch.float8_e4m3fnuz, + ScalarType.UINT16: torch.uint16, + ScalarType.UINT32: torch.uint32, + ScalarType.UINT64: torch.uint64, +} +_TORCH_DTYPE_TO_PTD: Dict[torch.dtype, ScalarType] = { + dtype: scalar_type for scalar_type, dtype in _PTD_TO_TORCH_DTYPE.items() +} +# PTD tensor layouts do not encode quantization parameters. +_QUANTIZED_DTYPES: Set[torch.dtype] = { + torch.qint8, + torch.quint8, + torch.qint32, + torch.quint4x2, + torch.quint2x4, +} + def _serialize_to_flatbuffer(flat_tensor: FlatTensor) -> Cord: """Serializes a FlatTensor to a flatbuffer and returns the serialized data.""" @@ -444,3 +486,114 @@ def deserialize_to_named_data_store_output( pte_data={}, external_data={name: data_payload.named_data}, ) + + +def save_ptd( + path: Union[str, os.PathLike[str]], tensor_map: Dict[str, torch.Tensor] +) -> None: + """Saves a dictionary of tensors to a PTD file. + + Args: + path: Path to the output PTD file. + tensor_map: Mapping from tensor names to supported strided CPU tensors. + """ + buffers: List[bytes] = [] + named_data: Dict[str, DataEntry] = {} + for name, tensor in tensor_map.items(): + if tensor.device.type != "cpu": + raise ValueError( + f"Tensor '{name}' must be on CPU, received {tensor.device}." + ) + if tensor.dtype not in _TORCH_DTYPE_TO_PTD or tensor.is_quantized: + raise ValueError(f"Unsupported tensor dtype {tensor.dtype} for PTD files.") + if tensor.layout != torch.strided: + raise ValueError(f"Tensor '{name}' must use a strided layout.") + + tensor_to_save = tensor.detach().resolve_neg() + if 0 in tensor_to_save.stride(): + tensor_to_save = tensor_to_save.clone(memory_format=torch.contiguous_format) + elif not ( + tensor_to_save.is_contiguous() + or tensor_to_save.is_contiguous(memory_format=torch.channels_last) + ): + tensor_to_save = tensor_to_save.contiguous() + if tensor_to_save.nbytes != tensor_to_save.untyped_storage().nbytes(): + tensor_to_save = tensor_to_save.clone(memory_format=torch.preserve_format) + tensor_layout = TensorLayout( + scalar_type=_TORCH_DTYPE_TO_PTD[tensor_to_save.dtype], + sizes=list(tensor_to_save.shape), + dim_order=list(dim_order_from_stride(tensor_to_save.stride())), + ) + + buffer_index = len(buffers) + buffers.append(bytes(tensor_to_save.untyped_storage())) + named_data[name] = DataEntry( + buffer_index=buffer_index, + alignment=1, + tensor_layout=tensor_layout, + ) + + data_payload = DataPayload(buffers=buffers, named_data=named_data) + serialized_data = FlatTensorSerializer().serialize(data_payload) + with open(path, "wb") as file: + file.write(bytes(serialized_data)) + + +def load_ptd(path: Union[str, os.PathLike[str]]) -> Dict[str, torch.Tensor]: + """Loads a dictionary of tensors from a PTD file. + + Args: + path: Path to the PTD file. + + Returns: + A mapping from tensor names to mutable CPU tensors. + """ + with open(path, "rb") as file: + data_payload = FlatTensorSerializer().deserialize(Cord(file.read())) + + tensor_map: Dict[str, torch.Tensor] = {} + for name, data_entry in data_payload.named_data.items(): + if not 0 <= data_entry.buffer_index < len(data_payload.buffers): + raise ValueError( + f"Tensor '{name}' references invalid buffer index " + f"{data_entry.buffer_index}." + ) + + tensor_layout = data_entry.tensor_layout + if tensor_layout is None: + raise ValueError(f"Named data '{name}' does not contain a tensor layout.") + + try: + dtype = _PTD_TO_TORCH_DTYPE[tensor_layout.scalar_type] + except KeyError as error: + raise ValueError( + f"Unsupported scalar type {tensor_layout.scalar_type} for tensor " + f"'{name}'." + ) from error + if dtype in _QUANTIZED_DTYPES: + raise ValueError(f"Unsupported tensor dtype {dtype} for PTD files.") + + strides = stride_from_dim_order(tensor_layout.sizes, tensor_layout.dim_order) + numel = math.prod(tensor_layout.sizes) + buffer = data_payload.buffers[data_entry.buffer_index] + expected_nbytes = ( + numel * torch.empty((), dtype=dtype, device="cpu").element_size() + ) + if len(buffer) < expected_nbytes: + raise ValueError( + f"Tensor '{name}' requires {expected_nbytes} bytes, but its buffer " + f"contains {len(buffer)} bytes." + ) + if numel == 0: + tensor = torch.empty_strided( + tensor_layout.sizes, strides, dtype=dtype, device="cpu" + ) + else: + tensor = torch.frombuffer( + bytearray(buffer), + dtype=dtype, + count=numel, + ).as_strided(tensor_layout.sizes, strides) + tensor_map[name] = tensor + + return tensor_map diff --git a/extension/flat_tensor/test/test_serialize.py b/extension/flat_tensor/test/test_serialize.py index 6ecd6911ac8..033a0dc3dd4 100644 --- a/extension/flat_tensor/test/test_serialize.py +++ b/extension/flat_tensor/test/test_serialize.py @@ -8,7 +8,9 @@ import dataclasses import math +import tempfile import unittest +from pathlib import Path from typing import Dict, List, Optional @@ -27,6 +29,7 @@ from executorch.exir.schema import ScalarType from executorch.exir.tensor_layout import TensorLayout +from executorch.extension.flat_tensor.serialize import load_ptd, save_ptd from executorch.extension.flat_tensor.serialize.serialize import ( _deserialize_to_flat_tensor, _FLATBUFFER_ALIGNMENT, @@ -280,6 +283,149 @@ def test_round_trip(self) -> None: TEST_DATA_PAYLOAD.named_data, deserialized_payload.named_data ) + def test_save_and_load_ptd(self) -> None: + weight = ( + torch.arange(48, dtype=torch.float32) + .reshape(2, 2, 3, 4) + .contiguous(memory_format=torch.channels_last)[1:] + ) + tensor_map = { + "weight": torch.nn.Parameter(weight), + "bias": torch.tensor([1, -2, 3], dtype=torch.int64), + "empty": torch.empty((0, 3), dtype=torch.float64), + "negative_view": torch._neg_view( + torch.tensor([1.0, -2.0], dtype=torch.float32) + ), + } + + with tempfile.TemporaryDirectory() as temp_dir: + path = Path(temp_dir) / "state.ptd" + save_ptd(path, tensor_map) + loaded_tensor_map = load_ptd(path) + + self.assertEqual(loaded_tensor_map.keys(), tensor_map.keys()) + for name, expected in tensor_map.items(): + with self.subTest(name=name): + actual = loaded_tensor_map[name] + self.assertEqual(actual.device.type, "cpu") + self.assertEqual(actual.dtype, expected.dtype) + self.assertEqual(actual.shape, expected.shape) + self.assertEqual(actual.stride(), expected.stride()) + torch.testing.assert_close(actual, expected) + + loaded_tensor_map["bias"].add_(1) + torch.testing.assert_close(loaded_tensor_map["bias"], tensor_map["bias"] + 1) + + def test_save_and_load_empty_ptd(self) -> None: + with tempfile.TemporaryDirectory() as temp_dir: + path = Path(temp_dir) / "empty.ptd" + save_ptd(path, {}) + self.assertEqual(load_ptd(path), {}) + + def test_save_and_load_additional_ptd_dtypes(self) -> None: + tensor_map = { + "bits16": torch.tensor([1, 2], dtype=torch.uint16).view(torch.bits16), + "float8_e4m3fn": torch.tensor([1, -2], dtype=torch.float8_e4m3fn), + "float8_e4m3fnuz": torch.tensor([1, -2], dtype=torch.float8_e4m3fnuz), + "float8_e5m2fnuz": torch.tensor([1, -2], dtype=torch.float8_e5m2fnuz), + "uint64": torch.tensor([1, 2**63], dtype=torch.uint64), + } + with tempfile.TemporaryDirectory() as temp_dir: + path = Path(temp_dir) / "additional_dtypes.ptd" + save_ptd(path, tensor_map) + loaded_tensor_map = load_ptd(path) + + for name, expected in tensor_map.items(): + with self.subTest(name=name): + actual = loaded_tensor_map[name] + self.assertEqual(actual.dtype, expected.dtype) + self.assertTrue( + torch.equal(actual.view(torch.uint8), expected.view(torch.uint8)) + ) + + def test_save_ptd_normalizes_unsupported_layouts(self) -> None: + tensor_map = { + "transposed": torch.arange(12, dtype=torch.float32).reshape(3, 4).t(), + "expanded": torch.arange(3, dtype=torch.float32).expand(4, 3), + "channels_last_3d": torch.arange(120, dtype=torch.float32) + .reshape(1, 2, 3, 4, 5) + .contiguous(memory_format=torch.channels_last_3d), + "empty_channels_last": torch.empty( + (2, 0, 3, 4), memory_format=torch.channels_last + ), + "empty_channels_last_3d": torch.empty( + (2, 0, 3, 4, 5), memory_format=torch.channels_last_3d + ), + } + with tempfile.TemporaryDirectory() as temp_dir: + path = Path(temp_dir) / "normalized_layouts.ptd" + save_ptd(path, tensor_map) + loaded_tensor_map = load_ptd(path) + + for name, expected in tensor_map.items(): + with self.subTest(name=name): + actual = loaded_tensor_map[name] + self.assertTrue(actual.is_contiguous()) + torch.testing.assert_close(actual, expected) + + def test_save_ptd_rejects_unsupported_tensors(self) -> None: + unsupported_tensors = { + "complex": ( + torch.tensor([1 + 2j], dtype=torch.complex64), + "torch.complex64", + ), + "quantized": ( + torch.quantize_per_tensor( + torch.tensor([1.0]), scale=0.5, zero_point=0, dtype=torch.qint8 + ), + "torch.qint8", + ), + "sparse": ( + torch.tensor([1.0]).to_sparse(), + "must use a strided layout", + ), + "meta": ( + torch.empty(1, device="meta"), + "must be on CPU", + ), + } + with tempfile.TemporaryDirectory() as temp_dir: + path = Path(temp_dir) / "unsupported.ptd" + for name, (tensor, error) in unsupported_tensors.items(): + with self.subTest(name=name), self.assertRaisesRegex(ValueError, error): + save_ptd(path, {name: tensor}) + + def test_load_ptd_rejects_invalid_tensor_data(self) -> None: + invalid_payloads = { + "opaque": DataPayload( + buffers=[b"data"], + named_data={ + "opaque": DataEntry(buffer_index=0, alignment=1, tensor_layout=None) + }, + ), + "short": DataPayload( + buffers=[b"\x00" * 4], + named_data={ + "short": DataEntry( + buffer_index=0, + alignment=1, + tensor_layout=TensorLayout( + scalar_type=ScalarType.FLOAT, + sizes=[2], + dim_order=[0], + ), + ) + }, + ), + } + with tempfile.TemporaryDirectory() as temp_dir: + path = Path(temp_dir) / "invalid.ptd" + for name, payload in invalid_payloads.items(): + serialized_data = FlatTensorSerializer().serialize(payload) + path.write_bytes(bytes(serialized_data)) + with self.subTest(name=name), self.assertRaises(ValueError): + load_ptd(path) + def test_deserialize_to_named_data_store_output(self) -> None: store = NamedDataStore() external_tag = "model"