(
config: _config.TrainConfig, init_rng: at.KeyArrayLike, mesh: jax.sharding.Mesh, *, resume: bool
)
| 83 | |
| 84 | @at.typecheck |
| 85 | def init_train_state( |
| 86 | config: _config.TrainConfig, init_rng: at.KeyArrayLike, mesh: jax.sharding.Mesh, *, resume: bool |
| 87 | ) -> tuple[training_utils.TrainState, Any]: |
| 88 | tx = _optimizer.create_optimizer(config.optimizer, config.lr_schedule, weight_decay_mask=None) |
| 89 | |
| 90 | def init(rng: at.KeyArrayLike, partial_params: at.Params | None = None) -> training_utils.TrainState: |
| 91 | rng, model_rng = jax.random.split(rng) |
| 92 | # initialize the model (and its parameters). |
| 93 | model = config.model.create(model_rng) |
| 94 | |
| 95 | # Merge the partial params into the model. |
| 96 | if partial_params is not None: |
| 97 | graphdef, state = nnx.split(model) |
| 98 | # This will produce an error if the partial params are not a subset of the state. |
| 99 | state.replace_by_pure_dict(partial_params) |
| 100 | model = nnx.merge(graphdef, state) |
| 101 | |
| 102 | params = nnx.state(model) |
| 103 | # Convert frozen params to bfloat16. |
| 104 | params = nnx_utils.state_map(params, config.freeze_filter, lambda p: p.replace(p.value.astype(jnp.bfloat16))) |
| 105 | |
| 106 | return training_utils.TrainState( |
| 107 | step=0, |
| 108 | params=params, |
| 109 | model_def=nnx.graphdef(model), |
| 110 | tx=tx, |
| 111 | opt_state=tx.init(params.filter(config.trainable_filter)), |
| 112 | ema_decay=config.ema_decay, |
| 113 | ema_params=None if config.ema_decay is None else params, |
| 114 | ) |
| 115 | |
| 116 | train_state_shape = jax.eval_shape(init, init_rng) |
| 117 | state_sharding = sharding.fsdp_sharding(train_state_shape, mesh, log=True) |
| 118 | |
| 119 | if resume: |
| 120 | return train_state_shape, state_sharding |
| 121 | |
| 122 | partial_params = _load_weights_and_validate(config.weight_loader, train_state_shape.params.to_pure_dict()) |
| 123 | replicated_sharding = jax.sharding.NamedSharding(mesh, jax.sharding.PartitionSpec()) |
| 124 | |
| 125 | # Initialize the train state and mix in the partial params. |
| 126 | train_state = jax.jit( |
| 127 | init, |
| 128 | donate_argnums=(1,), # donate the partial params buffer. |
| 129 | in_shardings=replicated_sharding, |
| 130 | out_shardings=state_sharding, |
| 131 | )(init_rng, partial_params) |
| 132 | |
| 133 | return train_state, state_sharding |
| 134 | |
| 135 | |
| 136 | @at.typecheck |
no test coverage detected