Skip to content

Commit a851995

Browse files
committed
fix(nnx): update lora regex word boundary and multi-host array resharding
1 parent d1a8cb1 commit a851995

3 files changed

Lines changed: 7 additions & 6 deletions

File tree

src/maxtext/configs/post_train/lora_module_path.yml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,7 @@ mistral: "decoder/layers/.*(attention/(query|key|value|out)|mlp/(wi_0|wi_1|wo))"
2121
deepseek2: "decoder/(dense_layers|moe_stack)/self_attention/(query|out|wkv_a|wkv_b)|decoder/(dense_layers|moe_stack)/(mlp|shared_experts)/(wi_0|wi_1|wo)"
2222
gemma2: "decoder/(scanned_blocks|layers_remainder|layers)/(self_attention_local|self_attention_global)/(query|key|value|out)|decoder/(scanned_blocks|layers_remainder|layers)/(mlp_local|mlp_global)/(wi_0|wi_1|wo)"
2323
gemma3: "decoder/(scanned_blocks|layers_remainder|layers)/.*(self_attention/(query|key|value|out)|mlp/(wi_0|wi_1|wo|gate|up|down))"
24-
gemma4: "decoder/((scanned_blocks|layers_remainder)/)?layers.*/.*(self_attention/(query|key|value|out)|mlp/.*(wi_0|wi_1|wo|shared_experts/(wi_0|wi_1|wo)))"
24+
gemma4: 'decoder/((scanned_blocks|layers_remainder)/)?layers.*/.*(self_attention/(query|key|value|out)\b|mlp/.*(wi_0|wi_1|wo|shared_experts/(wi_0|wi_1|wo)))'
2525
olmo3: "decoder/layers/.*(attention/(query|key|value|out)|mlp/(wi_0|wi_1|wo))"
2626
gpt3: "decoder/layers/(self_attention/(qkv_proj|out)|mlp/(wi|wo))"
2727

src/maxtext/utils/lora_utils.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -634,9 +634,8 @@ def _safe_reshard(var, sharding_spec):
634634
return var
635635
if not hasattr(val, "shape"):
636636
return var
637-
# make_array_from_callback natively constructs a globally sharded array
638-
# from the local host arrays, bypassing backend-specific device_put issues
639-
# on both Pathways and McJAX.
637+
if isinstance(val, jax.Array):
638+
return var
640639
resharded_val = jax.make_array_from_callback(val.shape, sharding_spec, lambda idx: val[idx])
641640
return var.replace(value=resharded_val)
642641

src/maxtext/utils/maxtext_utils.py

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1712,7 +1712,9 @@ def setup_initial_state(
17121712
# checkpoint didn't carry (or a params-only load didn't touch) keeps its real init
17131713
# value, so restore doesn't depend on knowing exactly what was saved.
17141714
state = jax.jit(
1715-
lambda: nnx.state(init_state_partial()), # Get state only, mapping to out_sharding structure
1715+
lambda: nnx.state(
1716+
init_state_partial(), nnx.Not(nnx.Intermediate)
1717+
), # Get state only, mapping to out_sharding structure
17161718
in_shardings=None,
17171719
out_shardings=state_mesh_shardings,
17181720
)()
@@ -1911,7 +1913,7 @@ def get_abstract_state_nnx(config, mesh, nnx_init_trainstate_fn, is_training=Tru
19111913
# which has no .shape and would raise AttributeError. We handle sharding
19121914
# ourselves via nnx_construct_named_sharding, so auto-assignment is not needed here.
19131915
abs_model = nnx.eval_shape(nnx_init_trainstate_fn)
1914-
_, abs_var_state = nnx.split(abs_model)
1916+
abs_var_state = nnx.state(abs_model, nnx.Not(nnx.Intermediate))
19151917
named_sharding_state = sharding.nnx_construct_named_sharding(abs_var_state, mesh)
19161918
abstract_state = jax.tree.map(
19171919
lambda a, s: jax.ShapeDtypeStruct(a.shape, a.dtype, sharding=s),

0 commit comments

Comments
 (0)