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,
)
| 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() |