(argv: Sequence[str])
| 195 | |
| 196 | |
| 197 | def main(argv: Sequence[str]) -> None: |
| 198 | jax.config.update("jax_default_prng_impl", "unsafe_rbg") |
| 199 | os.environ["LIBTPU_INIT_ARGS"] = ( |
| 200 | os.environ.get("LIBTPU_INIT_ARGS", "") + " --xla_tpu_spmd_rng_bit_generator_unsafe=true" |
| 201 | ) |
| 202 | print("Starting train_compile.py...", flush=True) |
| 203 | |
| 204 | # Parse and validate configuration |
| 205 | config = pyconfig.initialize(argv) |
| 206 | validate_config(config) |
| 207 | |
| 208 | # Create target mesh |
| 209 | topology_mesh = get_topology_mesh(config) |
| 210 | |
| 211 | # Print system information after building the compile topology to avoid |
| 212 | # prematurely initializing the backend. |
| 213 | max_utils.print_system_information() |
| 214 | |
| 215 | # Get shaped inputs |
| 216 | shaped_train_args, shaped_train_kwargs, state_mesh_shardings, model = get_shaped_inputs(topology_mesh, config) |
| 217 | |
| 218 | # Get data sharding |
| 219 | data_sharding = sharding.get_input_data_sharding(config, topology_mesh) |
| 220 | |
| 221 | # Get function to compile and shardings |
| 222 | func_to_compile, in_shard, out_shard, static_argnums, donate_argnums = ( |
| 223 | maxtext_utils.get_functional_train_with_signature( |
| 224 | train.train_step, data_sharding, state_mesh_shardings, model, config |
| 225 | ) |
| 226 | ) |
| 227 | |
| 228 | # print weights sharding info under debug sharding mode |
| 229 | if config.debug_sharding: |
| 230 | max_utils.print_non_trivial_mesh_axis(topology_mesh) |
| 231 | maxtext_utils.print_state_mesh_shardings_params(shaped_train_args[0], state_mesh_shardings, topology_mesh) |
| 232 | |
| 233 | # Compile |
| 234 | print("Jitting and compiling train step...", flush=True) |
| 235 | compiled = jit_and_compile( |
| 236 | func_to_compile, |
| 237 | shaped_train_args, |
| 238 | shaped_train_kwargs, |
| 239 | topology_mesh, |
| 240 | in_shard, |
| 241 | out_shard, |
| 242 | static_argnums, |
| 243 | donate_argnums, |
| 244 | nn_partitioning.axis_rules(config.logical_axis_rules), |
| 245 | ) |
| 246 | print("Jitting and compilation complete!", flush=True) |
| 247 | |
| 248 | # Serialize and save the compiled object |
| 249 | if config.compiled_trainstep_file != "": |
| 250 | print("Saving compiled object...") |
| 251 | save_compiled(compiled, config.compiled_trainstep_file) |
| 252 | print(f"Successfully saved compiled object as {config.compiled_trainstep_file}") |
| 253 | print("Finished train_compile.py successfully!", flush=True) |
| 254 | print(f"Cost analysis: {compiled.cost_analysis()}") |
nothing calls this directly
no test coverage detected