MCPcopy Create free account
hub / github.com/AIS-SNU/Smart-Infinity / _configure_optimizer

Method _configure_optimizer

deepspeed/runtime/engine.py:1156–1204  ·  view source on GitHub ↗
(self, client_optimizer, model_parameters, args)

Source from the content-addressed store, hash-verified

1154
1155 # Configure optimizer
1156 def _configure_optimizer(self, client_optimizer, model_parameters, args):
1157 if client_optimizer is not None:
1158 if isinstance(client_optimizer, tuple(self._supported_optims())):
1159 client_optimizer.param_groups[:] = [
1160 pg for pg in client_optimizer.param_groups if len(pg["params"]) != 0
1161 ]
1162 log_dist("Removing param_group that has no 'params' in the client Optimizer", ranks=[0])
1163
1164 basic_optimizer = client_optimizer
1165 log_dist('Using client Optimizer as basic optimizer', ranks=[0])
1166 else:
1167 basic_optimizer = client_optimizer(model_parameters)
1168 log_dist('Using client callable to create basic optimizer', ranks=[0])
1169
1170 if self.zero_use_cpu_optimizer() and not isinstance(basic_optimizer, deepspeed.ops.adam.DeepSpeedCPUAdam):
1171 if self.zero_force_ds_cpu_optimizer():
1172 msg = f'You are using ZeRO-Offload with a client provided optimizer ({type(basic_optimizer)}) which in most cases will yield poor performance. Please either use deepspeed.ops.adam.DeepSpeedCPUAdam or set an optimizer in your ds-config (https://www.deepspeed.ai/docs/config-json/#optimizer-parameters). If you really want to use a custom optimizer w. ZeRO-Offload and understand the performance impacts you can also set <"zero_force_ds_cpu_optimizer": false> in your configuration file.'
1173 raise ZeRORuntimeException(msg)
1174 else:
1175 basic_optimizer = self._configure_basic_optimizer(model_parameters)
1176 log_dist(f"Using DeepSpeed Optimizer param name {self.optimizer_name()} as basic optimizer", ranks=[0])
1177
1178 self._check_for_duplicates(basic_optimizer)
1179
1180 self.basic_optimizer = basic_optimizer
1181 log_dist("DeepSpeed Basic Optimizer = {}".format(basic_optimizer.__class__.__name__), ranks=[0])
1182
1183 optimizer_wrapper = self._do_optimizer_sanity_check(basic_optimizer)
1184
1185 if optimizer_wrapper == ZERO_OPTIMIZATION:
1186 self.optimizer = self._configure_zero_optimizer(basic_optimizer, args)
1187 elif optimizer_wrapper == AMP:
1188 amp_params = self.amp_params()
1189 log_dist(f"Initializing AMP with these params: {amp_params}", ranks=[0])
1190 model, self.optimizer = amp.initialize(self.module, basic_optimizer, **amp_params)
1191 self._set_client_model(model)
1192 self._broadcast_model()
1193 # TODO: maybe need to broadcast experts differently?
1194 elif optimizer_wrapper == FP16:
1195 self.optimizer = self._configure_fp16_optimizer(basic_optimizer)
1196 elif optimizer_wrapper == BFLOAT16:
1197 self.optimizer = self._configure_bf16_optimizer(basic_optimizer)
1198 else:
1199 self.optimizer = basic_optimizer
1200
1201 log_dist("DeepSpeed Final Optimizer = {}".format(self.optimizer_name()), ranks=[0])
1202
1203 self.compression_scheduler = self._configure_compression_scheduler()
1204 self.quantizer = self._configure_quantization()
1205
1206 def _configure_basic_optimizer(self, model_parameters):
1207 optimizer_parameters = self.optimizer_params()

Callers 1

__init__Method · 0.95

Calls 15

_supported_optimsMethod · 0.95
optimizer_nameMethod · 0.95
_check_for_duplicatesMethod · 0.95
amp_paramsMethod · 0.95
_set_client_modelMethod · 0.95
_broadcast_modelMethod · 0.95

Tested by

no test coverage detected