Files
truecluster/truecluster/protocol/tensors.py
T
2026-06-04 18:11:26 -05:00

89 lines
2.4 KiB
Python

from __future__ import annotations
from dataclasses import dataclass
from typing import Any
import numpy as np
import torch
_DTYPE_TO_TORCH = {
"float16": torch.float16,
"float32": torch.float32,
"float64": torch.float64,
"int64": torch.int64,
"int32": torch.int32,
"int16": torch.int16,
"int8": torch.int8,
"uint8": torch.uint8,
"bool": torch.bool,
}
_TORCH_TO_NUMPY = {
torch.float16: np.float16,
torch.float32: np.float32,
torch.float64: np.float64,
torch.int64: np.int64,
torch.int32: np.int32,
torch.int16: np.int16,
torch.int8: np.int8,
torch.uint8: np.uint8,
torch.bool: np.bool_,
}
@dataclass(frozen=True)
class SerializedTensor:
dtype: str
shape: list[int]
data: bytes
def to_payload(self) -> dict[str, Any]:
return {"dtype": self.dtype, "shape": self.shape, "data": self.data}
class TensorSerializationError(RuntimeError):
pass
def dtype_name(dtype: torch.dtype) -> str:
text = str(dtype)
if text.startswith("torch."):
return text.split(".", 1)[1]
return text
def serialize_tensor(tensor: torch.Tensor) -> dict[str, Any]:
"""Serialize a tensor into msgpack-friendly metadata + raw bytes.
Tensors are moved to CPU and made contiguous. bfloat16 is converted to float16
because NumPy has inconsistent bfloat16 support and the prototype currently
targets fp16 transfer.
"""
tensor = tensor.detach().cpu().contiguous()
if tensor.dtype == torch.bfloat16:
tensor = tensor.to(torch.float16)
if tensor.dtype not in _TORCH_TO_NUMPY:
raise TensorSerializationError(f"unsupported tensor dtype for serialization: {tensor.dtype}")
array = tensor.numpy()
return {
"dtype": dtype_name(tensor.dtype),
"shape": list(tensor.shape),
"data": array.tobytes(order="C"),
}
def deserialize_tensor(payload: dict[str, Any], device: str | torch.device | None = None) -> torch.Tensor:
dtype = payload["dtype"]
shape = list(payload["shape"])
data = payload["data"]
torch_dtype = _DTYPE_TO_TORCH.get(dtype)
if torch_dtype is None:
raise TensorSerializationError(f"unsupported tensor dtype for deserialization: {dtype}")
np_dtype = _TORCH_TO_NUMPY[torch_dtype]
array = np.frombuffer(data, dtype=np_dtype).copy().reshape(shape)
tensor = torch.from_numpy(array)
if device is not None:
tensor = tensor.to(device)
return tensor