more setitem kernel tests (#14748)

check where realize happened
This commit is contained in:
chenyu
2026-02-14 09:57:46 -05:00
committed by GitHub
parent 4ab51b55bd
commit 446909fb7a
+32 -1
View File
@@ -1,5 +1,5 @@
import unittest
from tinygrad import Tensor, TinyJit, Variable, dtypes, Device
from tinygrad import Tensor, TinyJit, Variable, dtypes, Device, GlobalCounters
from tinygrad.helpers import Context
import numpy as np
@@ -70,18 +70,44 @@ class TestSetitem(unittest.TestCase):
with self.assertRaises(RuntimeError): t[2:4] = Tensor([1, 2], dtype=dtypes.int)
def test_setitem_into_empty(self):
GlobalCounters.reset()
t = Tensor.empty(4)
self.assertEqual(GlobalCounters.kernel_count, 0)
t[1] = 5
self.assertEqual(GlobalCounters.kernel_count, 1)
t[1].realize()
self.assertEqual(GlobalCounters.kernel_count, 1)
t.realize()
self.assertEqual(GlobalCounters.kernel_count, 1)
self.assertEqual(t[1].item(), 5)
def test_setitem_into_tensor(self):
t = Tensor([1, 2, 3, 4]).realize()
GlobalCounters.reset()
self.assertEqual(GlobalCounters.kernel_count, 0)
t[1] = 5
self.assertEqual(GlobalCounters.kernel_count, 0)
t[1].realize()
self.assertEqual(GlobalCounters.kernel_count, 1)
t.realize()
self.assertEqual(GlobalCounters.kernel_count, 1)
self.assertListEqual(t.tolist(), [1, 5, 3, 4])
def test_setitem_into_cont(self):
t = Tensor.ones(4)
with self.assertRaises(RuntimeError): t[1] = 5
def test_setitem_into_const_alu(self):
# TODO: this is not consistent
GlobalCounters.reset()
t = Tensor.ones(4) + Tensor.ones(4)
self.assertEqual(GlobalCounters.kernel_count, 0)
t[1] = 5
self.assertEqual(GlobalCounters.kernel_count, 2)
t[1].realize()
self.assertEqual(GlobalCounters.kernel_count, 2)
t.realize()
self.assertEqual(GlobalCounters.kernel_count, 2)
self.assertListEqual(t.tolist(), [2, 5, 2, 2])
t = Tensor.ones(4) + Tensor.ones(4)
@@ -90,8 +116,13 @@ class TestSetitem(unittest.TestCase):
def test_setitem_into_arange(self):
# NOTE: arange has no real buffer, but assigning to it is fine
GlobalCounters.reset()
t = Tensor.arange(4)
self.assertEqual(GlobalCounters.kernel_count, 0)
t[1] = 5
self.assertEqual(GlobalCounters.kernel_count, 2)
t.realize()
self.assertEqual(GlobalCounters.kernel_count, 2)
self.assertListEqual(t.tolist(), [0, 5, 2, 3])
def test_setitem_chained_indexing(self):