(
model,
optimizer,
weights_only=False,
use_param_info_optim=True,
max_gradient_norm=None,
allow_lr_injection=False,
)
| 2126 | |
| 2127 | |
| 2128 | def _build( |
| 2129 | model, |
| 2130 | optimizer, |
| 2131 | weights_only=False, |
| 2132 | use_param_info_optim=True, |
| 2133 | max_gradient_norm=None, |
| 2134 | allow_lr_injection=False, |
| 2135 | ): |
| 2136 | param_to_device = _get_param_to_device(model) |
| 2137 | |
| 2138 | # Validate there are no duplicate params |
| 2139 | model.Validate() |
| 2140 | |
| 2141 | params = [] |
| 2142 | for param_info in model.GetOptimizationParamInfo(): |
| 2143 | if weights_only and param_info.blob not in model.weights: |
| 2144 | continue |
| 2145 | params.append(param_info) |
| 2146 | |
| 2147 | lr_multiplier = None |
| 2148 | if max_gradient_norm is not None: |
| 2149 | lr_multiplier = _calc_norm_ratio( |
| 2150 | model, |
| 2151 | params, |
| 2152 | "norm_clipped_grad_update", |
| 2153 | param_to_device, |
| 2154 | max_gradient_norm, |
| 2155 | ) |
| 2156 | |
| 2157 | if allow_lr_injection: |
| 2158 | if not model.net.BlobIsDefined(_LEARNING_RATE_INJECTION): |
| 2159 | lr_injection = model.param_init_net.ConstantFill( |
| 2160 | [], _LEARNING_RATE_INJECTION, shape=[1], value=1.0 |
| 2161 | ) |
| 2162 | else: |
| 2163 | lr_injection = _LEARNING_RATE_INJECTION |
| 2164 | |
| 2165 | if lr_multiplier is None: |
| 2166 | lr_multiplier = lr_injection |
| 2167 | else: |
| 2168 | lr_multiplier = model.net.Mul( |
| 2169 | [lr_multiplier, lr_injection], "lr_multiplier", broadcast=1 |
| 2170 | ) |
| 2171 | optimizer.add_lr_multiplier(lr_multiplier) |
| 2172 | |
| 2173 | for param_info in params: |
| 2174 | param_name = str(param_info.blob) |
| 2175 | device = get_param_device(param_name, param_info.grad, param_to_device) |
| 2176 | with core.DeviceScope(device): |
| 2177 | if param_info.optimizer and use_param_info_optim: |
| 2178 | param_info.optimizer(model.net, model.param_init_net, param_info) |
| 2179 | else: |
| 2180 | optimizer(model.net, model.param_init_net, param_info) |
| 2181 | return optimizer |
| 2182 | |
| 2183 | |
| 2184 | def add_weight_decay(model, weight_decay): |
no test coverage detected
searching dependent graphs…