@@ -23,7 +23,7 @@ class XlaShardingTest(unittest.TestCase):
2323
2424 class SimpleLinear (nn .Module ):
2525 def __init__ (self ):
26- super (SimpleLinear , self ).__init__ ()
26+ super (XlaShardingTest . SimpleLinear , self ).__init__ ()
2727 self .fc1 = nn .Linear (128 , 64 )
2828 self .relu = nn .ReLU ()
2929 self .fc2 = nn .Linear (64 , 1 )
@@ -47,6 +47,8 @@ def _get_mesh(self, mesh_shape, device_ids=None):
4747 assert len (device_ids ) == self .n_devices
4848 return xs .Mesh (device_ids , mesh_shape )
4949
50+ class BasicShardingTest (XlaShardingTest ):
51+
5052 def test_xla_sharded_tensor (self ):
5153 partition_spec = (0 , 1 )
5254 xt1 = torch .tensor ([[1 , 2 , 3 , 4 , 5 , 6 , 7 , 8 ]],
@@ -117,6 +119,15 @@ def test_deep_copy(self):
117119 torch_xla ._XLAC ._get_xla_sharding_spec (xt ),
118120 torch_xla ._XLAC ._get_xla_sharding_spec (xt2 ))
119121
122+ def test_clone (self ):
123+ xt = torch .randn (2 , 4 , 8 , 16 ).to (xm .xla_device ())
124+ xs .mark_sharding (xt , self ._get_mesh ((1 , 1 , 1 , self .n_devices )),
125+ (0 , 1 , 2 , 3 ))
126+ xt2 = xt .clone ()
127+ self .assertEqual (
128+ torch_xla ._XLAC ._get_xla_sharding_spec (xt ),
129+ torch_xla ._XLAC ._get_xla_sharding_spec (xt2 ))
130+
120131 def test_mark_step_with_sharding (self ):
121132 xt = torch .ones (2 , 2 ).to (xm .xla_device ())
122133 xs .mark_sharding (xt , self ._get_mesh ((1 , self .n_devices )), (0 , 1 ))
@@ -133,10 +144,11 @@ def test_optimizer_step_with_sharding(self):
133144 optimizer = optim .SGD (model .parameters (), lr = 0.1 )
134145 data = torch .randn (128 , 128 ).to (xm .xla_device ())
135146 target = torch .zeros (128 ).to (xm .xla_device ())
147+ loss_fn = nn .CrossEntropyLoss ()
136148 for i in range (5 ):
137149 optimizer .zero_grad ()
138150 output = model (data )
139- loss = nn . CrossEntropy (output , target )
151+ loss = loss_fn (output , target )
140152 loss .backward ()
141153 optimizer .step ()
142154 xm .mark_step ()
0 commit comments