From 5cf4ad2fb6a50e8464bfc645cf84a9f568c3befe Mon Sep 17 00:00:00 2001 From: nimlgen <138685161+nimlgen@users.noreply.github.com> Date: Thu, 23 Apr 2026 17:41:44 +0300 Subject: [PATCH] fix resolve param (#15889) --- test/backend/test_multitensor.py | 8 ++++++++ tinygrad/engine/realize.py | 2 +- 2 files changed, 9 insertions(+), 1 deletion(-) diff --git a/test/backend/test_multitensor.py b/test/backend/test_multitensor.py index 70e99ee8f1..38471b3842 100644 --- a/test/backend/test_multitensor.py +++ b/test/backend/test_multitensor.py @@ -276,6 +276,14 @@ class TestMultiTensor(unittest.TestCase): out = f(tt) assert out.item() == 1+2+3+4 + def test_multitensor_jit_input_reduce_shard_axis(self): + @TinyJit + def f(x): return x.sum(0).realize() + for _ in range(5): + tt = Tensor.ones(2, 64).contiguous().realize().shard((d1,d2), 0).realize() + out = f(tt) + np.testing.assert_allclose(out.numpy(), np.full(64, 2.0)) + def test_multitensor_inside_jit(self): @TinyJit def f(x): return (x.shard((d1,d2), 0)+1).contiguous().sum() diff --git a/tinygrad/engine/realize.py b/tinygrad/engine/realize.py index 4456ffa13d..eec31fa130 100644 --- a/tinygrad/engine/realize.py +++ b/tinygrad/engine/realize.py @@ -203,7 +203,7 @@ class ExecContext: jit: bool = False def _resolve(b:UOp, inputs:tuple[UOp, ...]) -> UOp: - if b.op is Ops.BUFFER_VIEW and b.src[0].op is Ops.PARAM: return b.replace(src=(inputs[b.src[0].arg], *b.src[1:])) + if b.op in (Ops.BUFFER_VIEW, Ops.MSELECT) and b.src[0].op is Ops.PARAM: return b.replace(src=(inputs[b.src[0].arg], *b.src[1:])) return inputs[b.arg] if b.op is Ops.PARAM else b def resolve_params(call:UOp, inputs:tuple[UOp, ...]) -> list[UOp]: return [_resolve(b, inputs) for b in call.src[1:] if b.op is not Ops.BIND]