MCPcopy Create free account
hub / github.com/THUDM/GLM / get_optimizer_param_groups

Function get_optimizer_param_groups

train_utils.py:159–172  ·  view source on GitHub ↗
(model)

Source from the content-addressed store, hash-verified

157
158
159def 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
175def get_optimizer(param_groups, args):

Callers 1

Tested by

no test coverage detected