Skip to content

Commit ebe8aba

Browse files
fix(nnx): Add backward-compatibility safe accessors for Flax NNX containers
1 parent bbcc873 commit ebe8aba

5 files changed

Lines changed: 46 additions & 16 deletions

File tree

src/maxtext/common/train_state_nnx.py

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -220,8 +220,16 @@ def split_for_checkpoint(state: nnx.State):
220220
params, rng_state, batch_stats, caches, intermediates, rest = nnx.split_state(
221221
state, nnx.Param, nnx.RngState, nnx.BatchStat, nnx.Cache, nnx.Intermediate, ...
222222
)
223-
optimizer = nnx.State({"optimizer": rest["optimizer"]}) if "optimizer" in rest else nnx.State({})
224-
custom = nnx.State({"model": rest["model"]}) if "model" in rest else nnx.State({})
223+
optimizer = (
224+
nnx.State({"optimizer": rest["optimizer"].to_pure_dict() if hasattr(rest["optimizer"], "to_pure_dict") else rest["optimizer"]})
225+
if "optimizer" in rest
226+
else nnx.State({})
227+
)
228+
custom = (
229+
nnx.State({"model": rest["model"].to_pure_dict() if hasattr(rest["model"], "to_pure_dict") else rest["model"]})
230+
if "model" in rest
231+
else nnx.State({})
232+
)
225233
linen_state = nnx.merge_state(params, optimizer)
226234
aux = nnx.merge_state(rng_state, batch_stats, custom)
227235
ephemeral = nnx.merge_state(caches, intermediates)

src/maxtext/layers/embeddings.py

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -153,7 +153,10 @@ def __call__(self, inputs: Array, model_mode: str = MODEL_MODE_TRAIN) -> Array:
153153
raise ValueError("Input type must be an integer or unsigned integer.")
154154

155155
embedding = jnp.asarray(
156-
_maybe_move_embedding_to_device(self.embedding.get_value(), self.config),
156+
_maybe_move_embedding_to_device(
157+
self.embedding.raw_value if hasattr(self.embedding, "raw_value") else getattr(self.embedding, "value", self.embedding),
158+
self.config,
159+
),
157160
self.dtype,
158161
)
159162

@@ -197,7 +200,7 @@ def attend(self, query: Array, out_sharding: NamedSharding | None = None) -> Arr
197200
Commonly used for weight-sharing between embeddings and logit transform
198201
in NLP models.
199202
"""
200-
embedding = self.embedding.get_value()
203+
embedding = self.embedding.raw_value if hasattr(self.embedding, "raw_value") else getattr(self.embedding, "value", self.embedding)
201204
attend_dtype = self.attend_dtype if self.attend_dtype is not None else self.dtype
202205
return attend_on_embedding(query, embedding, attend_dtype, self.config, out_sharding)
203206

src/maxtext/trainers/pre_train/train.py

Lines changed: 6 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -78,9 +78,8 @@
7878
def get_first_step(model, state):
7979
if isinstance(model, nn.Module):
8080
return int(state.step)
81-
if hasattr(state, "inner_state"): # DiLoCoTrainState (NNX DiLoCo): step is the optimizer step var
82-
return int(state.step.get_value())
83-
return int(state.optimizer.step.get_value())
81+
step_var = state.step if hasattr(state, "inner_state") else state.optimizer.step
82+
return int(step_var.value if hasattr(step_var, "value") else step_var.raw_value if hasattr(step_var, "raw_value") else step_var)
8483

8584

8685
# -----------------------------------------------------------------------------
@@ -420,7 +419,10 @@ def train_step(model, config, state_mesh_shardings, params_shardings, state, dat
420419
nnx.update(state.model, curr_params)
421420

422421
def diff_wrapper(curr_params, custom_params, rest, config, data):
423-
local_model = nnx.merge(model_graphdef, curr_params, custom_params, rest, copy=True)
422+
try:
423+
local_model = nnx.merge(model_graphdef, curr_params, custom_params, rest, copy=True)
424+
except TypeError:
425+
local_model = nnx.merge(model_graphdef, curr_params, custom_params, rest)
424426
loss, aux = loss_fn(local_model, config, data, None, None, is_train=True)
425427
_, _, _, new_rest = nnx.split(local_model, nnx.Param, custom_param_filter, ...)
426428
return loss, (aux, new_rest)

src/maxtext/utils/maxtext_utils_nnx.py

Lines changed: 19 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -180,12 +180,22 @@ def create_sharded_state():
180180
return nnx.merge(graphdef, sharded_state)
181181

182182

183+
def _get_val(x):
184+
if hasattr(x, "raw_value"):
185+
return x.raw_value
186+
if hasattr(x, "value"):
187+
return x.value
188+
if hasattr(x, "get_value"):
189+
return x.get_value()
190+
return x
191+
192+
183193
def nnx_ensure_scan_leading_axis(tree, length):
184194
"""Broadcasts scalar-like variables to have a leading scan axis."""
185195

186196
def _op(x):
187197
is_var = isinstance(x, nnx.Variable)
188-
val = x.get_value() if is_var else x
198+
val = _get_val(x)
189199
if hasattr(val, "shape") and len(val.shape) == 0:
190200
new_val = jax.numpy.broadcast_to(val, (length,))
191201
return x.replace(value=new_val) if is_var else new_val
@@ -214,6 +224,8 @@ def nnx_update_sharding_meta(variable, transform_fn):
214224
updates[key] = P(*transformed) if isinstance(val, P) else tuple(transformed)
215225

216226
if updates:
227+
if "raw_value" not in updates and "value" not in updates:
228+
updates["raw_value"] = _get_val(variable)
217229
return variable.replace(**updates)
218230
return variable
219231

@@ -234,14 +246,14 @@ def remove_fn(l):
234246
removed = axis_name in l
235247
if removed:
236248
l.remove(axis_name)
237-
if len(l) > x.get_value().ndim:
249+
if len(l) > _get_val(x).ndim:
238250
if removed:
239251
raise ValueError(
240-
f"Sharding names {l} still exceed value rank {x.get_value().ndim} after removing scan axis "
252+
f"Sharding names {l} still exceed value rank {_get_val(x).ndim} after removing scan axis "
241253
f"{axis_name!r}; the partition metadata is inconsistent."
242254
)
243255
raise ValueError(
244-
f"Scan axis {axis_name!r} not found in sharding names {l} for a rank-{x.get_value().ndim} value; "
256+
f"Scan axis {axis_name!r} not found in sharding names {l} for a rank-{_get_val(x).ndim} value; "
245257
"the partition metadata is inconsistent."
246258
)
247259
return l
@@ -267,17 +279,17 @@ def _op(x):
267279
axis_name = x.get_metadata().get(nnx.PARTITION_NAME, name)
268280
target = x.get_metadata().get("param_scan_axis", pos)
269281

270-
val = x.get_value()
282+
val = _get_val(x)
271283
if target != 0 and hasattr(val, "ndim") and val.ndim > target:
272284
x = x.replace(value=jnp.moveaxis(val, 0, target))
273285

274286
def add_fn(l):
275287
if axis_name not in l:
276-
while len(l) < x.get_value().ndim - 1:
288+
while len(l) < _get_val(x).ndim - 1:
277289
l.append(None)
278290
l.insert(target, axis_name)
279291
else:
280-
while len(l) < x.get_value().ndim:
292+
while len(l) < _get_val(x).ndim:
281293
l.append(None)
282294
return l
283295

src/maxtext/utils/sharding.py

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -178,7 +178,12 @@ def remove_size_one_mesh_axis(spec, mesh):
178178

179179
def get_nnx_var_named_sharding_with_scan_axis(v: nnx.Variable, mesh) -> nnx.Variable:
180180
"""Compute NamedSharding for an NNX variable, correctly handling the scan axis."""
181-
val = v.get_value()
181+
if hasattr(v, "raw_value"):
182+
val = v.raw_value
183+
elif hasattr(v, "value"):
184+
val = v.value
185+
else:
186+
val = v
182187
if not hasattr(val, "shape"):
183188
# `val` is either truly leafless (e.g. optax MaskedNode) or a composite
184189
# pytree of tensors (e.g. AQT QTensor on serve-mode quantized variables).

0 commit comments

Comments
 (0)