MCPcopy Create free account
hub / github.com/deepspeedai/DeepSpeed / __init__

Method __init__

deepspeed/runtime/engine.py:337–614  ·  view source on GitHub ↗
(self,
                 args,
                 model,
                 optimizer=None,
                 model_parameters=None,
                 training_data=None,
                 lr_scheduler=None,
                 mpu=None,
                 dist_init_required=None,
                 collate_fn=None,
                 config=None,
                 config_class=None,
                 mesh_device=None,
                 dont_change_device=False)

Source from the content-addressed store, hash-verified

335 r"""DeepSpeed engine for training."""
336
337 def __init__(self,
338 args,
339 model,
340 optimizer=None,
341 model_parameters=None,
342 training_data=None,
343 lr_scheduler=None,
344 mpu=None,
345 dist_init_required=None,
346 collate_fn=None,
347 config=None,
348 config_class=None,
349 mesh_device=None,
350 dont_change_device=False):
351 super(DeepSpeedEngine, self).__init__()
352 self.dont_change_device = dont_change_device
353 self.client_optimizer = optimizer
354 self.client_lr_scheduler = lr_scheduler
355 self.training_data = training_data
356 self.collate_fn = collate_fn
357 self.mpu = mpu
358 self.all_to_all_group = None
359 self.data_parallel_group = None
360 self.global_steps = 0
361 self.global_samples = 0
362 self.micro_steps = 0
363 # Unmanaged mode: backward() calls since the last step(), used to advance global_samples.
364 self._unmanaged_backward_count = 0
365 self.skipped_steps = 0
366 self.gradient_average = config_class.gradient_allreduce_op == GRADIENT_ALLREDUCE_OP_MEAN
367 self.warn_unscaled_loss = True
368 self.config = config
369 self._config = config_class
370 self.loaded_checkpoint_mp_world_size = None
371 self.loaded_checkpoint_dp_world_size = None
372 self.enable_backward_allreduce = True
373 self.inside_no_sync_ctxt = False
374 self.progressive_layer_drop = None
375 self.eigenvalue = None
376 self.block_eigenvalue = None
377 self.gas_boundary_ctr = 0
378 self.dist_backend = get_accelerator().communication_backend_name()
379 self.has_moe_layers = False
380 self.num_experts = []
381 self.gate_modules = []
382 self.moe_layers = []
383 self._step_applied = False
384 self._global_grad_norm = None
385 self.use_ds_comm = False # False --> Use torch.dist, True --> Use ds.comm backend.
386 self.checkpoint_engine = None
387 self.optimizer = None
388 self.basic_optimizer = None
389 self.lr_scheduler = None
390
391 self._is_gradient_accumulation_boundary = None
392 self.scale_wrt_gas = None
393 self.losses = None
394 self.mesh_device = mesh_device

Callers

nothing calls this directly

Calls 15

_do_args_sanity_checkMethod · 0.95
_do_sanity_checkMethod · 0.95
log_levelMethod · 0.95
autotp_sizeMethod · 0.95
memory_breakdownMethod · 0.95
elasticity_enabledMethod · 0.95
_set_distributed_varsMethod · 0.95

Tested by

no test coverage detected