(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)
| 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 |
nothing calls this directly
no test coverage detected