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

Method __init__

axlearn/common/trainer.py:246–381  ·  view source on GitHub ↗
(
        self,
        cfg: Config,
        *,
        parent: Optional[Module],
        devices: Optional[np.ndarray] = None,
    )

Source from the content-addressed store, hash-verified

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

Callers

nothing calls this directly

Calls 15

_maybe_record_eventMethod · 0.95
_step_logMethod · 0.95
meshMethod · 0.95
maybe_instantiateFunction · 0.90
maybe_set_configFunction · 0.90
ParameterSpecClass · 0.90
TrainerStateClass · 0.85
flattenMethod · 0.80
_add_childMethod · 0.80
joinMethod · 0.80
mapMethod · 0.80

Tested by

no test coverage detected