few more dtype cast convinience methods (#11480)

This commit is contained in:
chenyu
2025-08-02 15:47:09 -04:00
committed by GitHub
parent e22e5da9a5
commit 66be747908
2 changed files with 20 additions and 0 deletions
+15
View File
@@ -1,4 +1,5 @@
import unittest
from tinygrad.tensor import Tensor
from tinygrad.dtype import dtypes, DType, ImageDType, PtrDType, to_dtype
class TestImageDType(unittest.TestCase):
@@ -38,5 +39,19 @@ class TestToDtype(unittest.TestCase):
self.assertIsInstance(res, DType)
self.assertEqual(res, dtypes.int32)
class TestCastConvenienceMethod(unittest.TestCase):
def test_method(self):
for input_dtype in (dtypes.float, dtypes.int):
t = Tensor([1, 2], dtype=input_dtype)
self.assertEqual(t.dtype, input_dtype)
self.assertEqual(t.bool().dtype, dtypes.bool)
self.assertEqual(t.short().dtype, dtypes.short)
self.assertEqual(t.int().dtype, dtypes.int)
self.assertEqual(t.long().dtype, dtypes.long)
self.assertEqual(t.half().dtype, dtypes.half)
self.assertEqual(t.bfloat16().dtype, dtypes.bfloat16)
self.assertEqual(t.float().dtype, dtypes.float)
self.assertEqual(t.double().dtype, dtypes.double)
if __name__ == "__main__":
unittest.main()
+5
View File
@@ -4316,6 +4316,11 @@ class Tensor(MathTrait):
"""
return self.cast(dtypes.bool)
def bfloat16(self) -> Tensor: return self.cast(dtypes.bfloat16)
def double(self) -> Tensor: return self.cast(dtypes.double)
def long(self) -> Tensor: return self.cast(dtypes.long)
def short(self) -> Tensor: return self.cast(dtypes.short)
# *** image Tensor function replacements ***
def image_dot(self, w:Tensor, dtype:DTypeLike|None=None) -> Tensor: