Repository navigation
nnx.Optimizer doesn't respect extra sharding axis added by nnx.scan/nnx.vmap #5112
Description
Activity
The issue seems to be with the following function in flax:
flax/flax/nnx/training/optimizer.py
Line 49 in 697f4e5
def to_opt_state(tree): It tries to apply sharding of the model's tensor which contains sharding metadata for original, non-vmapped Variable.
Generally, I think that for vmapped arrays, we should automatically add additional axis to its sharding metadata, shouldn't we? Because rn we have this discrepancy between sharding metadata VariableMetadata and actual sharding of the array which lead to issue like this one.
Thanks for stress testing flax in these cases, @qGentry !
Let me see what happens inside.Reacted by Filipp FisinSeems like extending nnx.vmap args with
transform_metadatasolves this issue.@nnx.vmap( in_axes=(0,), out_axes=0, transform_metadata={ nnx.spmd.PARTITION_NAME: None, } )I wonder if it should happen automatically or mentioned in documentation here
https://flax.readthedocs.io/en/latest/nnx_basics.html#scan-over-layersI was able to reproduce the issue, and I think the problem comes from how sharding metadata is handled when NNX transforms add new axes. I think vmap, scan, and pmap should always insert a default sharding metadata entry ({PARTITION_NAME: None}) so the extra axis they introduce is explicitly marked as unsharded.
Also, probably the sharding add/remove helper functions should be idempotent and tolerant of missing axes, so sharding_names stays aligned with the actual array rank after a lift. So, when nnx.Optimizer rewraps parameters and reapplies eager sharding; it would no longer try to shard the new leading axis, which should avoid the dimension-0 divisibility error.
@vfdev-5 Any thoughts or suggestions?
To summarize what's been said so far, the core issue is that when we
nnx.vmapover a function that returns a Param, we currently get back a Param that has the same sharding metadata rather than one with an extra None in it. You have to explicitly add thennx.PARTITION_NAMEkey intransform_metadataif you want the returned Param to have the proper sharding. This is already documented in the nnx guides https://flax.readthedocs.io/en/stable/guides/transforms.html#axis-metadata. But it is pretty unintuitive.One way to get around this is to use Jax's explicit sharding instead:
mesh = jax.make_mesh((2, 4), ("a", "b"), axis_types=(AxisType.Explicit, AxisType.Explicit)) jax.set_mesh(mesh) class Model(nnx.Module): def __init__(self, num_layers, rngs: nnx.Rngs): @nnx.split_rngs(splits=num_layers) @nnx.vmap(in_axes=(0,), out_axes=0) def create_linear(rngs: nnx.Rngs): return nnx.Param( jnp.ones((16, 16), out_sharding=P("a", "b")) ) self.linears = create_linear(rngs=rngs)
This works just fine with the Optimizer case above.
@samanklesaria Good point on the metadata gap. I’ve put together a fix that makes vmap/scan/pmap insert a default {PARTITION_NAME: None} so sharding metadata tracks the extra axis automatically, and made the axis add/remove helpers idempotent. This avoids having to set transform_metadata manually or use explicit sharding for the Optimizer case. I’ll open a PR could you take a look?
@mohsinm-dev that seems reasonable, but I'll have to check with @cgarciae , who will be on break until next week.
@samanklesaria, cool. let me know after you will discuss with @cgarciae then we can discuss and work on it.
I've also see quite unexpected results even with the use of transform_metadata:
from flax import nnx import jax import jax.numpy as jnp spmd_mesh = jax.make_mesh((2, 1, 4), ("layers", "fsdp", "tensor")) jax.set_mesh(spmd_mesh) class Block(nnx.Module): def __init__(self, in_features, out_features, rngs): self.dense1 = nnx.Linear( in_features=in_features, out_features=out_features, kernel_init=nnx.with_partitioning( nnx.initializers.lecun_normal(), ("fsdp", "tensor"), ), rngs=rngs, ) self.dense2 = nnx.Linear( in_features=out_features, out_features=in_features, kernel_init=nnx.with_partitioning( nnx.initializers.lecun_normal(), ("tensor", "fsdp"), ), rngs=rngs, ) self.relu = nnx.relu def __call__(self, x): x = self.dense1(x) x = self.relu(x) x = self.dense2(x) return x class Model(nnx.Module): def __init__(self, in_features, out_features, n_layers, rngs): @nnx.split_rngs(splits=n_layers) @nnx.vmap(in_axes=0, out_axes=0, transform_metadata={nnx.PARTITION_NAME: "layers"}) def get_blocks(rngs): return Block(in_features=in_features, out_features=out_features, rngs=rngs) self.blocks = get_blocks(rngs) @jax.jit def get_model(): model = Model(in_features=16, out_features=16, n_layers=8, rngs=nnx.Rngs(params=0)) return model model = get_model() print(f"dense1 sharding_names: {model.blocks.dense1.kernel.sharding_names}") print(f"dense1 sharding: {model.blocks.dense1.kernel.sharding}")
output:
dense1 sharding_names: ('layers', 'fsdp', 'tensor') dense1 sharding: NamedSharding(mesh=Mesh('layers': 2, 'fsdp': 1, 'tensor': 4, axis_types=(Auto, Auto, Auto)), spec=PartitionSpec(None, None, 'tensor'), memory_kind=device)For some reason, dense1 is not sharded over the leading axis even though "layers" is in sharding_names
@qGentry The
transform_metadata={nnx.PARTITION_NAME: "layers"}keyword argument is about transforming Variable metadata. In the model you present, there is no variable metadata:nnx.with_partitioning(nnx.initializers.lecun_normal(), ("tensor", "fsdp"))shards your kernels, but doesn't encode that sharding within flax as metadata. Here's how you would write the code using flax metadata:import jax jax.config.update('jax_num_cpu_devices', 8) from flax import nnx import jax.numpy as jnp spmd_mesh = jax.make_mesh((2, 2, 2), ("layers", "fsdp", "tensor"), axis_types=(jax.sharding.AxisType.Auto,) * 3) jax.set_mesh(spmd_mesh) def make_model(rngs): return nnx.Linear( in_features=16, out_features=16, kernel_metadata={'sharding_names': ('fsdp', 'tensor')}, rngs=rngs) def test_problem(): @nnx.split_rngs(splits=2) @nnx.vmap(in_axes=0, out_axes=0, transform_metadata={nnx.PARTITION_NAME: "layers"}) def get_blocks(rngs): return make_model(rngs) result = get_blocks(nnx.Rngs(0)) print(result.kernel.sharding)
The
transform_metadataargument doesn't do any sharding itself. It's just there to notify flax about how to update the metadata about variables that are sharded when you transform them.@samanklesaria
nnx.with_partitioningdoes result in metadata being added to the Variable.@qGentry I'm starting to incline to a yes after thinking about this, feel free to send a PR.
Hi folks, me again.
I keep playing around with nnx and seems like nnx.Optimizer, when creating optimizer state for models intended to be used with 'scan', use sharding information for original, non-stacked tensor, without taking into the account extra dimension added by vmap.
Repro script:
Output: