Returns compiled XLA programs for the given trainer. Args: trainer_config: The trainer config. topology: A string representing the TPU topology, e.g., "v4-8". Must be a key in USER_FACING_NAME_TO_SYSTEM_CHARACTERISTICS. If None, use CPU devices. t
(
trainer_config: SpmdTrainer.Config,
*,
topology: str,
topology_num_slices: int = 1,
compiler_options: Optional[Dict[str, Union[str, bool]]] = None,
)
| 366 | |
| 367 | |
| 368 | def compile_trainer_programs( |
| 369 | trainer_config: SpmdTrainer.Config, |
| 370 | *, |
| 371 | topology: str, |
| 372 | topology_num_slices: int = 1, |
| 373 | compiler_options: Optional[Dict[str, Union[str, bool]]] = None, |
| 374 | ) -> Dict[str, jax.stages.Compiled]: |
| 375 | """Returns compiled XLA programs for the given trainer. |
| 376 | |
| 377 | Args: |
| 378 | trainer_config: The trainer config. |
| 379 | topology: A string representing the TPU topology, e.g., "v4-8". Must be a key in |
| 380 | USER_FACING_NAME_TO_SYSTEM_CHARACTERISTICS. |
| 381 | If None, use CPU devices. |
| 382 | topology_num_slices: The number of TPU slices. |
| 383 | compiler_options: Options to pass to XLA. See `compiler_options.py` for examples. |
| 384 | |
| 385 | Returns: |
| 386 | A dict containing the following programs: |
| 387 | * "train_step": a program to run a single training step. |
| 388 | |
| 389 | Raises: |
| 390 | NotImplementedError: if `topology` is not in USER_FACING_NAME_TO_SYSTEM_CHARACTERISTICS. |
| 391 | """ |
| 392 | topology_devices, devices_per_slice = get_devices_for_topology(topology, topology_num_slices) |
| 393 | |
| 394 | cfg = trainer_config.clone(dir="NOT_USED") |
| 395 | cfg.mesh_axis_names = cfg.mesh_axis_names or ("data", "model") |
| 396 | # Use a default mesh_shape if None or REQUIRED. |
| 397 | cfg.mesh_shape = cfg.mesh_shape or [len(topology_devices)] + [1] * ( |
| 398 | len(cfg.mesh_axis_names) - 1 |
| 399 | ) |
| 400 | |
| 401 | topology_devices, mesh_shape = reshape_devices( |
| 402 | devices=topology_devices, |
| 403 | mesh_shape=cfg.mesh_shape, |
| 404 | devices_per_slice=devices_per_slice, |
| 405 | num_slices=topology_num_slices, |
| 406 | ) |
| 407 | cfg.mesh_shape = mesh_shape |
| 408 | |
| 409 | trainer: SpmdTrainer = cfg.instantiate(parent=None, devices=topology_devices) |
| 410 | compiled_train_step = trainer.compile_train_step(compiler_options=compiler_options) |
| 411 | return {"train_step": compiled_train_step} |
| 412 | |
| 413 | |
| 414 | def compile_inference_programs( |