MCPcopy Create free account
hub / github.com/apple/axlearn / should_update_with_optimizers

Method should_update_with_optimizers

axlearn/common/learner.py:234–243  ·  view source on GitHub ↗

Returns whether each parameter should be updated with the optimizers. Args: model_params: A nested structure with OptParams as leaf nodes. Returns: A nested dict with the same structure as `model_params` with boolean leaf values.

(self, model_params: Nested[OptParam])

Source from the content-addressed store, hash-verified

232 register_per_param_settings(
233 update_types, description="learner_update_type", path=self.path()
234 )
235 state = dict(
236 optimizer=self.optimizer.init(self._get_optimizer_model_params(model_params)),
237 )
238 if self.config.ema.decay is not None:
239 state["ema"] = self.ema.init(model_params)
240 return state
241
242 def _update_types(self, tree: dict) -> dict:
243 cfg = self.config
244 return jax.tree.map(
245 lambda path: match_regex_rules(
246 path, rules=cfg.update_rules, default_value=UpdateType.ALL_UPDATES

Callers 6

forward_and_backwardMethod · 0.95
read_per_param_settingsFunction · 0.45
_train_stepMethod · 0.45

Calls 2

_update_typesMethod · 0.95
mapMethod · 0.80

Tested by 1

read_per_param_settingsFunction · 0.36