forked from tinygrad/tinygrad
use truncate in onnx read_int64 [pr] (#11720)
This commit is contained in:
@@ -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("<f", self.read(4))[0]
|
||||
def read_packed_floats(self) -> 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()
|
||||
|
||||
Reference in New Issue
Block a user