(
self,
cfg: Config,
*,
parent: Optional[Module],
devices: Optional[np.ndarray] = None,
)
| 244 | # and compilation process every train step and rely on compilation cache to prevent |
| 245 | # excessive recompilations. Note: this could introduce overhead to training due to |
| 246 | # pre-compilation checks (such as sharding check) that increases the step time for some |
| 247 | # models. Note that this cache is always disabled at steps when xsc is enabled. |
| 248 | # Defaults to None which is interpreted as True. |
| 249 | cache_compiled_train_step: Optional[bool] = None |
| 250 | |
| 251 | # Log the loss value every n steps. Defaults to None which is interpreted as every |
| 252 | # 100 steps. |
| 253 | log_every_n_steps: Optional[int] = None |
| 254 | |
| 255 | def __init__( |
| 256 | self, |
| 257 | cfg: Config, |
| 258 | *, |
| 259 | parent: Optional[Module], |
| 260 | devices: Optional[np.ndarray] = None, |
| 261 | ): |
| 262 | super().__init__(cfg, parent=parent) |
| 263 | cfg = self.config |
| 264 | |
| 265 | if not cfg.prune_empty_state_updates: |
| 266 | raise ValueError( |
| 267 | "Setting prune_empty_state_updates to False is no longer supported.\n" |
| 268 | "The config option will be removed in a future AXLearn version." |
| 269 | ) |
| 270 | |
| 271 | self._step: int = None |
| 272 | self._trainer_state: TrainerState = None |
| 273 | self._jit_train_step: jax.stages.Wrapped = None |
| 274 | self._watchdog_stopping = None |
| 275 | self._watchdog_thread = None |
| 276 | self._device_monitor = maybe_instantiate(cfg.device_monitor) |
| 277 | self._recorder = maybe_instantiate(cfg.recorder) |
| 278 | self._is_initialized: bool = False |
| 279 | self._maybe_record_event(measurement.Event.START_ACCELERATOR_INIT) |
| 280 | |
| 281 | if cfg.model.dtype is None: |
| 282 | raise ValueError(f"dtype must be explicitly specified for {self.path()}.model") |
| 283 | if cfg.model.param_init is None: |
| 284 | cfg.model.param_init = DefaultInitializer.default_config() |
| 285 | logging.info( |
| 286 | "model.param_init is not specified. Default to DefaultInitializer: %s", |
| 287 | cfg.model.param_init, |
| 288 | ) |
| 289 | |
| 290 | self._per_param_train_dtype = maybe_instantiate( |
| 291 | canonicalize_per_param_dtype(cfg.train_dtype) |
| 292 | ) |
| 293 | |
| 294 | # Create the device mesh. |
| 295 | if devices is None: |
| 296 | self._step_log( |
| 297 | "Devices: global=%s local=%s %s", |
| 298 | jax.device_count(), |
| 299 | jax.local_device_count(), |
| 300 | [device.platform for device in jax.local_devices()], |
| 301 | ) |
| 302 | else: |
| 303 | local_devices = [d for d in devices.flatten() if d.process_index == jax.process_index()] |
nothing calls this directly
no test coverage detected