(model)
| 157 | |
| 158 | |
| 159 | def get_optimizer_param_groups(model): |
| 160 | # Build parameter groups (weight decay and non-decay). |
| 161 | while isinstance(model, (LocalDDP, TorchDDP, FP16_Module)): |
| 162 | model = model.module |
| 163 | param_groups = glm_get_params_for_weight_decay_optimization(model) |
| 164 | |
| 165 | # Add model parallel attribute if it is not set. |
| 166 | for param_group in param_groups: |
| 167 | # print('## param_group', len(param_group['params'])) |
| 168 | for param in param_group['params']: |
| 169 | if not hasattr(param, 'model_parallel'): |
| 170 | param.model_parallel = False |
| 171 | |
| 172 | return param_groups |
| 173 | |
| 174 | |
| 175 | def get_optimizer(param_groups, args): |
no test coverage detected