@@ -44,7 +44,7 @@ def test_to_local_requires_grad(self):
4444 tensor = torch .randn (100_000 , 88 , requires_grad = True )
4545
4646 # Create XLAShardedTensor
47- sharded_tensor = XLAShardedTensor (tensor , mesh , [Shard (0 )])
47+ sharded_tensor = XLAShardedTensor (tensor , mesh , [Shard (0 )], requires_grad = tensor . requires_grad )
4848
4949 # Verify requires_grad is set
5050 self .assertTrue (sharded_tensor .requires_grad )
@@ -70,7 +70,7 @@ def test_to_local_grad_independence(self):
7070 mesh = DeviceMesh ("xla" , list (range (world_size )))
7171
7272 tensor = torch .randn (100_000 , 88 , requires_grad = True )
73- sharded_tensor = XLAShardedTensor (tensor , mesh , [Shard (0 )])
73+ sharded_tensor = XLAShardedTensor (tensor , mesh , [Shard (0 )], requires_grad = tensor . requires_grad )
7474
7575 # Create gradients
7676 res = sharded_tensor .sum ()
@@ -95,7 +95,7 @@ def test_to_local_grad_none_handling(self):
9595 mesh = DeviceMesh ("xla" , list (range (world_size )))
9696
9797 tensor = torch .randn (100_000 , 88 , requires_grad = True )
98- sharded_tensor = XLAShardedTensor (tensor , mesh , [Shard (0 )])
98+ sharded_tensor = XLAShardedTensor (tensor , mesh , [Shard (0 )], requires_grad = tensor . requires_grad )
9999
100100 # Don't do backward pass, so grad remains None
101101 self .assertIsNone (sharded_tensor .grad )
0 commit comments