MCPcopy Create free account
hub / github.com/apple/axlearn / compile_train_step

Method compile_train_step

axlearn/common/trainer.py:1223–1266  ·  view source on GitHub ↗

Produce a lowered and compiled training step. Args: trainer_state: The global trainer state (or state specs). If None, infer from self.trainer_state_specs. input_batch: An input batch (or specs for the global input batch). If None, att

(
        self,
        *,
        trainer_state: Optional[TrainerState] = None,
        input_batch: Optional[dict[str, Any]] = None,
        compiler_options: Optional[dict[str, Union[str, bool]]] = None,
    )

Source from the content-addressed store, hash-verified

1221 in_shardings=(
1222 self._trainer_state_partition_specs,
1223 self._train_step_input_partition_specs(),
1224 ),
1225 out_shardings=(
1226 self._trainer_state_partition_specs,
1227 dict(
1228 summaries=None,
1229 loss=None,
1230 aux=None,
1231 ),
1232 ),
1233 donate_argnums=(0,), # donate the state
1234 )
1235
1236 def compile_train_step(
1237 self,
1238 *,
1239 trainer_state: Optional[TrainerState] = None,
1240 input_batch: Optional[dict[str, Any]] = None,
1241 compiler_options: Optional[dict[str, Union[str, bool]]] = None,
1242 ) -> jax.stages.Compiled:
1243 """Produce a lowered and compiled training step.
1244
1245 Args:
1246 trainer_state: The global trainer state (or state specs).
1247 If None, infer from self.trainer_state_specs.
1248 input_batch: An input batch (or specs for the global input batch).
1249 If None, attempt to infer from the (host-local) input element spec.
1250 compiler_options: Options passed to the XLA compiler, selectively overwriting
1251 any settings already provided by environment variables for this compilation.
1252
1253 Returns:
1254 A compiled training step, with signature matching self._pjit_train_step's return.
1255 """
1256 with self.mesh(), self._context_manager():
1257 if trainer_state is None:
1258 # Do not run init(), which requires real devices.
1259 trainer_state = jax.tree.map(
1260 lambda spec: jax.ShapeDtypeStruct(shape=spec.shape, dtype=spec.dtype),
1261 self.trainer_state_specs,
1262 )
1263 if input_batch is None:
1264 # Infer global input batch shapes from input element spec.
1265 host_batch = self.input.element_spec()
1266 if "input_dispatcher" in self.input.children:
1267 host_batch = self.input.input_dispatcher.logical_to_physical_shapes(host_batch)
1268 input_batch = host_to_global_specs(
1269 host_batch, partition=self._train_step_input_partition_specs()

Callers 3

compile_trainer_programsFunction · 0.80

Calls 10

meshMethod · 0.95
_pjit_train_stepMethod · 0.95
host_to_global_specsFunction · 0.90
aot_model_analysisFunction · 0.85
mapMethod · 0.80
lowerMethod · 0.80
compileMethod · 0.80
element_specMethod · 0.45

Tested by 1