From b67345caa3fa2cdcbe8e1669a12ae3ec643bef1a Mon Sep 17 00:00:00 2001 From: chenyu Date: Mon, 18 Aug 2025 17:49:35 -0700 Subject: [PATCH] use truncate in onnx read_int64 [pr] (#11720) --- tinygrad/frontend/onnx.py | 6 ++---- 1 file changed, 2 insertions(+), 4 deletions(-) diff --git a/tinygrad/frontend/onnx.py b/tinygrad/frontend/onnx.py index bcead3994d..d8a3cade76 100644 --- a/tinygrad/frontend/onnx.py +++ b/tinygrad/frontend/onnx.py @@ -5,7 +5,7 @@ from io import BufferedReader from tinygrad.nn.state import TensorIO from tinygrad.tensor import Tensor, _broadcast_shape, ReductionStr from tinygrad.helpers import getenv, DEBUG, all_same, prod, flatten, make_tuple, argsort, is_numpy_ndarray, get_single_element, polyN -from tinygrad.dtype import DType, ConstType, dtypes, _from_np_dtype +from tinygrad.dtype import DType, ConstType, dtypes, _from_np_dtype, truncate from tinygrad.device import is_dtype_supported, Device # ***** protobuf definitions ****** @@ -105,9 +105,7 @@ class PBBufferedReader(BufferedReader): def read_bytes(self) -> Tensor: return self.read_delimited(use_tensor=True) def read_float(self) -> float: return struct.unpack(" Tensor: return self.read_delimited(use_tensor=True) - def read_int64(self) -> int: - val = self.decode_varint() - return val - 2**64 if val & (1 << 63) else val + def read_int64(self) -> int: return truncate[dtypes.int64](self.decode_varint()) def read_packed_int64s(self) -> list[int]: total_bytes_len = self.decode_varint() old_pos = self.tell()