import torch from truecluster.protocol.tensors import deserialize_tensor, serialize_tensor def test_tensor_roundtrip_float16(): x = torch.randn(2, 3).half() y = deserialize_tensor(serialize_tensor(x)) assert y.dtype == torch.float16 assert y.shape == x.shape assert torch.equal(x, y)