(self, client_optimizer, model_parameters, args)
| 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() |
no test coverage detected