MCPcopy Create free account
hub / github.com/Physical-Intelligence/openpi / init_train_state

Function init_train_state

scripts/train.py:85–133  ·  view source on GitHub ↗
(
    config: _config.TrainConfig, init_rng: at.KeyArrayLike, mesh: jax.sharding.Mesh, *, resume: bool
)

Source from the content-addressed store, hash-verified

83
84@at.typecheck
85def 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

Callers 1

mainFunction · 0.85

Calls 1

Tested by

no test coverage detected