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

Function main

scripts/train.py:194–276  ·  view source on GitHub ↗
(config: _config.TrainConfig)

Source from the content-addressed store, hash-verified

192
193
194def main(config: _config.TrainConfig):
195 init_logging()
196 logging.info(f"Running on: {platform.node()}")
197
198 if config.batch_size % jax.device_count() != 0:
199 raise ValueError(
200 f"Batch size {config.batch_size} must be divisible by the number of devices {jax.device_count()}."
201 )
202
203 jax.config.update("jax_compilation_cache_dir", str(epath.Path("~/.cache/jax").expanduser()))
204
205 rng = jax.random.key(config.seed)
206 train_rng, init_rng = jax.random.split(rng)
207
208 mesh = sharding.make_mesh(config.fsdp_devices)
209 data_sharding = jax.sharding.NamedSharding(mesh, jax.sharding.PartitionSpec(sharding.DATA_AXIS))
210 replicated_sharding = jax.sharding.NamedSharding(mesh, jax.sharding.PartitionSpec())
211
212 checkpoint_manager, resuming = _checkpoints.initialize_checkpoint_dir(
213 config.checkpoint_dir,
214 keep_period=config.keep_period,
215 overwrite=config.overwrite,
216 resume=config.resume,
217 )
218 init_wandb(config, resuming=resuming, enabled=config.wandb_enabled)
219
220 data_loader = _data_loader.create_data_loader(
221 config,
222 sharding=data_sharding,
223 shuffle=True,
224 )
225 data_iter = iter(data_loader)
226 batch = next(data_iter)
227 logging.info(f"Initialized data loader:\n{training_utils.array_tree_to_info(batch)}")
228
229 # Log images from first batch to sanity check.
230 images_to_log = [
231 wandb.Image(np.concatenate([np.array(img[i]) for img in batch[0].images.values()], axis=1))
232 for i in range(min(5, len(next(iter(batch[0].images.values())))))
233 ]
234 wandb.log({"camera_views": images_to_log}, step=0)
235
236 train_state, train_state_sharding = init_train_state(config, init_rng, mesh, resume=resuming)
237 jax.block_until_ready(train_state)
238 logging.info(f"Initialized train state:\n{training_utils.array_tree_to_info(train_state.params)}")
239
240 if resuming:
241 train_state = _checkpoints.restore_state(checkpoint_manager, train_state, data_loader)
242
243 ptrain_step = jax.jit(
244 functools.partial(train_step, config),
245 in_shardings=(replicated_sharding, train_state_sharding, data_sharding),
246 out_shardings=(train_state_sharding, replicated_sharding),
247 donate_argnums=(1,),
248 )
249
250 start_step = int(train_state.step)
251 pbar = tqdm.tqdm(

Callers 1

train.pyFile · 0.70

Calls 4

init_train_stateFunction · 0.85
updateMethod · 0.80
init_loggingFunction · 0.70
init_wandbFunction · 0.70

Tested by

no test coverage detected