From 557134e1c7c1a5c09bbdb5aab3d9dfaa5f4d0519 Mon Sep 17 00:00:00 2001 From: chenyu Date: Thu, 12 Feb 2026 09:08:16 -0500 Subject: [PATCH] model/test fix that failed with WEBGPU=1 DEBUG=2 (#14706) --- extra/models/resnet.py | 2 +- test/null/test_tensor.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/extra/models/resnet.py b/extra/models/resnet.py index 016f1d0759..40662f5d0c 100644 --- a/extra/models/resnet.py +++ b/extra/models/resnet.py @@ -150,7 +150,7 @@ class ResNet: continue # Skip FC if transfer learning if 'bn' not in k and 'downsample' not in k: assert obj.shape == dat.shape, (k, obj.shape, dat.shape) - obj.assign(dat.to(obj.device).reshape(obj.shape)) + obj.assign(dat.to(obj.device).cast(obj.dtype).reshape(obj.shape)) ResNet18 = lambda num_classes=1000: ResNet(18, num_classes=num_classes) ResNet34 = lambda num_classes=1000: ResNet(34, num_classes=num_classes) diff --git a/test/null/test_tensor.py b/test/null/test_tensor.py index bccf304866..4bf3d68f06 100644 --- a/test/null/test_tensor.py +++ b/test/null/test_tensor.py @@ -113,13 +113,13 @@ class TestIdxUpcast(unittest.TestCase): @unittest.skipIf(is_dtype_supported(dtypes.long), "int64 is supported") def test_int64_unsupported_overflow_sym(self): - with self.assertRaises(KeyError): + with self.assertRaises((KeyError, RuntimeError)): self.do_op_then_assert(dtypes.long, 2048, 2048, UOp.variable("dim3", 1, 2048).bind(32)) @unittest.skipIf(is_dtype_supported(dtypes.long), "int64 is supported") @unittest.expectedFailure # bug in gpu dims limiting def test_int64_unsupported_overflow(self): - with self.assertRaises(KeyError): + with self.assertRaises((KeyError, RuntimeError)): self.do_op_then_assert(dtypes.long, 2048, 2048, 2048) @unittest.skip("This is kept for reference, it requires large memory to run")