Skip to content

Commit 6822d11

Browse files
SurbhiJainUSCGoogle-ML-Automation
authored andcommitted
Fix linting issue in stitch_checkpoint.py
PiperOrigin-RevId: 953501460
1 parent 99fcce5 commit 6822d11

1 file changed

Lines changed: 1 addition & 1 deletion

File tree

src/maxtext/experimental/omni_poc/utils/stitch_checkpoint.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -161,7 +161,7 @@ def stitch_and_save_checkpoints(
161161
init_params = nnx.state(model, nnx.Param)
162162
else:
163163
model = model_creation_utils.from_config(config, jax.devices())
164-
_, _, init_params = maxtext_utils.init_initial_state(model, None, config, is_training=False, init_rng=init_rng)
164+
_, _, init_params = maxtext_utils.init_initial_state(model, None, config, is_training=False, key=init_rng)
165165

166166
# Convert to pure pytree for easier processing
167167
is_nnx = isinstance(init_params, nnx.State)

0 commit comments

Comments
 (0)