89 lines
2.4 KiB
Python
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
|