MCPcopy Create free account
hub / github.com/pytorch/pytorch / _build

Function _build

caffe2/python/optimizer.py:2128–2181  ·  view source on GitHub ↗
(
    model,
    optimizer,
    weights_only=False,
    use_param_info_optim=True,
    max_gradient_norm=None,
    allow_lr_injection=False,
)

Source from the content-addressed store, hash-verified

2126
2127
2128def _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
2184def add_weight_decay(model, weight_decay):

Callers 14

add_weight_decayFunction · 0.85
build_sgdFunction · 0.85
build_fp16_sgdFunction · 0.85
build_ftrlFunction · 0.85
build_gftrlFunction · 0.85
build_adagradFunction · 0.85
build_wngradFunction · 0.85
build_stormFunction · 0.85
build_adadeltaFunction · 0.85
build_adamFunction · 0.85
build_decay_adagradFunction · 0.85

Calls 9

_get_param_to_deviceFunction · 0.85
_calc_norm_ratioFunction · 0.85
get_param_deviceFunction · 0.85
ValidateMethod · 0.80
BlobIsDefinedMethod · 0.80
add_lr_multiplierMethod · 0.80
optimizerMethod · 0.80
appendMethod · 0.45

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…