MCPcopy Create free account
hub / github.com/AI-Hypercomputer/maxtext / main

Function main

src/MaxText/train_compile.py:197–265  ·  view source on GitHub ↗
(argv: Sequence[str])

Source from the content-addressed store, hash-verified

195
196
197def 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()}")

Callers

nothing calls this directly

Calls 7

get_topology_meshFunction · 0.85
get_shaped_inputsFunction · 0.85
jit_and_compileFunction · 0.85
save_compiledFunction · 0.85
updateMethod · 0.80
validate_configFunction · 0.70
initializeMethod · 0.45

Tested by

no test coverage detected