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

Method __init__

deepspeed/runtime/engine.py:238–497  ·  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

236 r"""DeepSpeed engine for training."""
237
238 def __init__(self,
239 args,
240 model,
241 optimizer=None,
242 model_parameters=None,
243 training_data=None,
244 lr_scheduler=None,
245 mpu=None,
246 dist_init_required=None,
247 collate_fn=None,
248 config=None,
249 config_class=None,
250 mesh_device=None,
251 dont_change_device=False):
252 super(DeepSpeedEngine, self).__init__()
253 self.dont_change_device = dont_change_device
254 self.client_optimizer = optimizer
255 self.client_lr_scheduler = lr_scheduler
256 self.training_data = training_data
257 self.collate_fn = collate_fn
258 self.mpu = mpu
259 self.all_to_all_group = None
260 self.data_parallel_group = None
261 self.global_steps = 0
262 self.global_samples = 0
263 self.micro_steps = 0
264 self.skipped_steps = 0
265 self.gradient_average = True
266 self.warn_unscaled_loss = True
267 self.config = config
268 self._config = config_class
269 self.loaded_checkpoint_mp_world_size = None
270 self.loaded_checkpoint_dp_world_size = None
271 self.enable_backward_allreduce = True
272 self.inside_no_sync_ctxt = False
273 self.progressive_layer_drop = None
274 self.eigenvalue = None
275 self.block_eigenvalue = None
276 self.gas_boundary_ctr = 0
277 self.dist_backend = get_accelerator().communication_backend_name()
278 self.has_moe_layers = False
279 self.num_experts = []
280 self.gate_modules = []
281 self.moe_layers = []
282 self._step_applied = False
283 self._global_grad_norm = None
284 self.use_ds_comm = False # False --> Use torch.dist, True --> Use ds.comm backend.
285 self.checkpoint_engine = None
286 self.optimizer = None
287 self.basic_optimizer = None
288 self.lr_scheduler = None
289
290 self._is_gradient_accumulation_boundary = None
291 self.scale_wrt_gas = None
292 self.losses = None
293 self.mesh_device = mesh_device
294 self._autoep_folding_spec = None
295 self._autoep_folding_group_handles = None

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