(
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,
dont_change_device=False,
)
| 182 | r"""DeepSpeed engine for training.""" |
| 183 | |
| 184 | def __init__( |
| 185 | self, |
| 186 | args, |
| 187 | model, |
| 188 | optimizer=None, |
| 189 | model_parameters=None, |
| 190 | training_data=None, |
| 191 | lr_scheduler=None, |
| 192 | mpu=None, |
| 193 | dist_init_required=None, |
| 194 | collate_fn=None, |
| 195 | config=None, |
| 196 | config_class=None, |
| 197 | dont_change_device=False, |
| 198 | ): |
| 199 | super(DeepSpeedEngine, self).__init__() |
| 200 | self.dont_change_device = dont_change_device |
| 201 | self.client_optimizer = optimizer |
| 202 | self.client_lr_scheduler = lr_scheduler |
| 203 | self.training_data = training_data |
| 204 | self.collate_fn = collate_fn |
| 205 | self.mpu = mpu |
| 206 | self.data_parallel_group = None |
| 207 | self.global_steps = 0 |
| 208 | self.global_samples = 0 |
| 209 | self.micro_steps = 0 |
| 210 | self.skipped_steps = 0 |
| 211 | self.gradient_average = True |
| 212 | self.warn_unscaled_loss = True |
| 213 | self.config = config |
| 214 | self._config = config_class |
| 215 | self.loaded_checkpoint_mp_world_size = None |
| 216 | self.loaded_checkpoint_dp_world_size = None |
| 217 | self.enable_backward_allreduce = True |
| 218 | self.progressive_layer_drop = None |
| 219 | self.eigenvalue = None |
| 220 | self.block_eigenvalue = None |
| 221 | self.gas_boundary_ctr = 0 |
| 222 | self.dist_backend = get_accelerator().communication_backend_name() |
| 223 | self.has_moe_layers = False |
| 224 | self.num_experts = [] |
| 225 | self.gate_modules = [] |
| 226 | self.moe_layers = [] |
| 227 | self._step_applied = False |
| 228 | self._global_grad_norm = None |
| 229 | self.use_ds_comm = False # False --> Use torch.dist, True --> Use ds.comm backend. |
| 230 | |
| 231 | self.checkpoint_engine = None |
| 232 | |
| 233 | self._is_gradient_accumulation_boundary = None |
| 234 | self.scale_wrt_gas = None |
| 235 | self.losses = [] |
| 236 | |
| 237 | # for debug purposes - can then debug print: debug_get_module_name(module) |
| 238 | debug_extract_module_and_param_names(model) |
| 239 | |
| 240 | # needed for zero_to_fp32 weights reconstruction to remap nameless data to state_dict |
| 241 | self.param_names = {param: name for name, param in model.named_parameters()} |
nothing calls this directly
no test coverage detected