@@ -301,57 +301,35 @@ def __init__(self, rngs):
301301 assert jax .tree .leaves (abs_state )[0 ].sharding .is_equivalent_to (
302302 NamedSharding (mesh , P (None , 'model' )), ndim = 2 )
303303
304- def test_explicit_sharding (self ):
305- mesh = jax .make_mesh (
306- (2 , 2 ),
307- ('row' , 'col' ),
308- axis_types = (jax .sharding .AxisType .Auto , jax .sharding .AxisType .Explicit ),
309- )
310- v = nnx .Variable (
311- jnp .ones ((4 , 4 )),
312- sharding_names = ('row' , 'col' ),
313- mesh = mesh ,
314- )
315- self .assertEqual (v .sharding .mesh , mesh )
316- self .assertEqual (
317- v .sharding .spec ,
318- P ('row' , 'col' ),
319- )
304+ @parameterized .parameters ('auto' , 'explicit' , 'mixed' )
305+ def test_sharding_axis_types (self , mode ):
306+ if mode == 'auto' :
307+ axis_types = (jax .sharding .AxisType .Auto , jax .sharding .AxisType .Auto )
308+ elif mode == 'explicit' :
309+ axis_types = (jax .sharding .AxisType .Explicit , jax .sharding .AxisType .Explicit )
310+ else :
311+ axis_types = (jax .sharding .AxisType .Auto , jax .sharding .AxisType .Explicit )
320312
321- def test_explicit_sharding_disable_jit (self ):
322313 mesh = jax .make_mesh (
323314 (2 , 2 ),
324315 ('row' , 'col' ),
325- axis_types = ( jax . sharding . AxisType . Auto , jax . sharding . AxisType . Explicit ) ,
316+ axis_types = axis_types ,
326317 )
327- with jax .disable_jit (True ):
318+ if mode == 'mixed' :
319+ with self .assertRaises (ValueError ):
320+ nnx .Variable (
321+ jnp .ones ((4 , 4 )),
322+ sharding_names = ('row' , 'col' ),
323+ mesh = mesh ,
324+ )
325+ else :
328326 v = nnx .Variable (
329327 jnp .ones ((4 , 4 )),
330328 sharding_names = ('row' , 'col' ),
331329 mesh = mesh ,
332330 )
333- self .assertEqual (v .sharding .mesh , mesh )
334- self .assertEqual (
335- v .sharding .spec ,
336- P ('row' , 'col' ),
337- )
338-
339- def test_explicit_sharding_mesh_context (self ):
340- mesh = jax .make_mesh (
341- (2 , 2 ),
342- ('row' , 'col' ),
343- axis_types = (jax .sharding .AxisType .Auto , jax .sharding .AxisType .Explicit ),
344- )
345- with jax .set_mesh (mesh ):
346- v = nnx .Variable (
347- jnp .ones ((4 , 4 )),
348- sharding_names = ('row' , 'col' ),
349- )
350- self .assertEqual (v .sharding .mesh , mesh )
351- self .assertEqual (
352- v .sharding .spec ,
353- P ('row' , 'col' ),
354- )
331+ self .assertEqual (v .sharding .mesh , mesh )
332+ self .assertEqual (v .sharding .spec , P ('row' , 'col' ))
355333
356334def has_sharding_spec (array ):
357335 sharding = array .sharding
0 commit comments