2828import drjax
2929from flax import nnx
3030from flax import struct
31- from flax .training import train_state
3231import jax
3332import jax .numpy as jnp
3433from jaxtyping import Array , Int32 , Key , PyTree , UInt32
35- import optax
36-
34+ from maxtext .common .train_state_nnx import TrainStateNNX
3735from maxtext .configs import pyconfig
36+ import optax
3837
3938Batch = Any
4039Params = PyTree
@@ -157,8 +156,12 @@ def add_diloco_dim(x):
157156 # For NNX, model params (Param variables only) live under abstract_state.model;
158157 # for Linen under abstract_state.params.
159158 if config .pure_nnx :
160- model_params = abstract_state .model .filter (nnx .Param )
161- model_params_sharding = state_mesh_shardings .model .filter (nnx .Param )
159+ _ , model_params , _ = nnx .split (abstract_state .model , nnx .Param , ...)
160+ model_params = model_params .to_pure_dict ()
161+ _ , model_params_sharding , _ = nnx .split (
162+ state_mesh_shardings .model , nnx .Param , ...
163+ )
164+ model_params_sharding = model_params_sharding .to_pure_dict ()
162165 else :
163166 model_params = abstract_state .params
164167 model_params_sharding = state_mesh_shardings .params
@@ -216,7 +219,11 @@ def init_diloco_state() -> tuple[DiLoCoTrainState, PyTree]:
216219 # Outer state retains a single copy of the model parameters and optimizer state.
217220 # For NNX, model params (Param variables only) live under state.model;
218221 # for Linen under state.params.
219- outer_params = state .model .filter (nnx .Param ) if config .pure_nnx else state .params
222+ if config .pure_nnx :
223+ _ , outer_params , _ = nnx .split (state .model , nnx .Param , ...)
224+ outer_params = outer_params .to_pure_dict ()
225+ else :
226+ outer_params = state .params
220227 outer_opt_state = outer_optimizer .init (outer_params )
221228 outer_opt_state_sharding = jax .tree_util .tree_map (lambda x : x .sharding , outer_opt_state )
222229 # For NNX, the step counter lives at state.optimizer.step; for Linen at state.step.
@@ -258,9 +265,13 @@ def synchronize(state):
258265 # state (since last synchronization).
259266 broadcast_outer_params = drjax .broadcast (state .params , mesh = mesh )
260267 # For NNX, model Param vars live under inner_state.model; for Linen under inner_state.params.
261- inner_model_params = (
262- nnx .filter_state (state .inner_state .model , nnx .Param ) if config .pure_nnx else state .inner_state .params
263- )
268+ if config .pure_nnx :
269+ _ , inner_model_params , _ = nnx .split (
270+ state .inner_state .model , nnx .Param , ...
271+ )
272+ inner_model_params = inner_model_params .to_pure_dict ()
273+ else :
274+ inner_model_params = state .inner_state .params
264275 model_delta = jax .tree .map (lambda x , y : y - x , inner_model_params , broadcast_outer_params )
265276 # Treat the average delta as the outer optimizer's gradient and apply to
266277 # the global (outer) model params.
@@ -273,15 +284,40 @@ def synchronize(state):
273284 if config .pure_nnx :
274285 # For NNX: merge new Param vars back with the non-Param model vars (e.g. RNG state).
275286 def replace_nnx_model_params (s , new_params ):
276- non_param_model = nnx .filter_state (s .model , nnx .Not (nnx .Param ))
277- new_model = nnx .merge_state (non_param_model , new_params )
278- # Assign via __setitem__ so nested States are stored as plain dicts (matching
279- # nnx.state()'s pytree structure). The dict-literal constructor keeps them as
280- # State objects, which makes jax.lax.cond see mismatched pytree structures.
281- result = type (s )({})
282- result ["model" ] = new_model
283- result ["optimizer" ] = s ["optimizer" ]
284- return result
287+ s_model = s ["model" ] if hasattr (s , "keys" ) else s .model
288+ s_opt = s ["optimizer" ] if hasattr (s , "keys" ) else s .optimizer
289+
290+ graphdef , _ , non_param_state = nnx .split (s_model , nnx .Param , ...)
291+ new_model = nnx .merge (graphdef , new_params , non_param_state )
292+
293+ if type (s_model ).__name__ == "State" :
294+ new_model = nnx .state (new_model )
295+ elif isinstance (s_model , dict ):
296+ new_model = nnx .to_pure_dict (new_model )
297+
298+ if hasattr (s , "keys" ):
299+ # Replace "model" leaves by path, keeping s's treedef. Picking by position
300+ # (leaves[N:]) breaks if a key sorts before "model"; reconstructing via
301+ # type(s)({...}) breaks the lax.cond match — nnx.State recursive-wraps.
302+ leaves_with_paths , treedef = jax .tree_util .tree_flatten_with_path (s )
303+ new_model_iter = iter (jax .tree_util .tree_leaves (new_model ))
304+
305+ def _is_model_leaf (path ):
306+ if not path :
307+ return False
308+ k = path [0 ]
309+ return (
310+ getattr (k , "key" , None ) == "model"
311+ or getattr (k , "name" , None ) == "model"
312+ )
313+
314+ new_leaves = [
315+ next (new_model_iter ) if _is_model_leaf (p ) else leaf
316+ for p , leaf in leaves_with_paths
317+ ]
318+ return jax .tree_util .tree_unflatten (treedef , new_leaves )
319+ else :
320+ return TrainStateNNX (new_model , s_opt )
285321
286322 new_inner_state = drjax .map_fn (
287323 lambda s : replace_nnx_model_params (s , new_outer_params ),
0 commit comments