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

Method update

axlearn/common/learner.py:502–552  ·  view source on GitHub ↗

Computes `model_params` updates from `update`. Args: updates: The updates to potentially transform and then apply. Returns: The updated model parameters. The learner state updates will be placed in the output collection's 'state_update' section.

(self, updates: Updates)

Source from the content-addressed store, hash-verified

500 cfg = self.config
501 learner_tree = self._learner_tree(params=model_params)
502 register_per_param_settings(learner_tree, description="learner_rule", path=self.path())
503 learner_state = {}
504 for name in cfg.learners.keys():
505 # Whether each parameter should apply the sub learner.
506 should_apply = jax.tree.map(
507 lambda learner_name, n=name: learner_name == n,
508 learner_tree,
509 )
510 # Mask model params.
511 sub_learner_model_params = mask_tree(
512 tree=model_params, keep=should_apply, mask_value=optax.MaskedNode()
513 )
514 # Call sub learner initialization.
515 sub_learner_state = getattr(self, name).init(model_params=sub_learner_model_params)
516 # Sub-learner's state.
517 learner_state[name] = sub_learner_state
518 if self.config.ema.decay is not None:
519 learner_state["ema"] = self.ema.init(model_params)
520 return learner_state
521
522 def update(self, updates: Updates) -> Nested[Tensor]:
523 """Computes `model_params` updates from `update`.
524
525 Args:
526 updates: The updates to potentially transform and then apply.
527
528 Returns:
529 The updated model parameters. The learner state updates will be placed in the output
530 collection's 'state_update' section.
531 """
532 cfg = self.config
533
534 updated_model_params = jax.tree.map(jnp.zeros_like, updates.param_values())
535
536 for name in cfg.learners.keys():
537 # Whether each parameter/state should apply the sub learner.
538 def should_apply(tree: Nested[Any]) -> Nested[bool]:
539 return jax.tree.map(
540 # pylint: disable-next=cell-var-from-loop
541 lambda learner_name, n=name: learner_name == n,
542 self._learner_tree(tree),
543 )
544
545 # See the docstring of `learner_test.CompositeLearnerTest.test_learner_masking`
546 # for a more detailed explanation of the masking behavior we are mimicking here
547 # for backwards compatibility.
548 sub_learner_updates = updates.mask(should_apply, fields=("inplace_updates",))
549 sub_learner_updates = sub_learner_updates.mask(
550 # pylint: disable-next=cell-var-from-loop
551 lambda _: should_apply(updates.opt_params),
552 fields=("opt_params", "delta_updates"),
553 )
554 sub_learner_updated_model_params = getattr(self, name).update(sub_learner_updates)
555 updated_model_params = jax.tree.map(

Callers 1

forward_and_backwardMethod · 0.95

Calls 7

mapMethod · 0.80
param_valuesMethod · 0.80
keysMethod · 0.80
replaceMethod · 0.80
maskMethod · 0.45
updateMethod · 0.45
add_state_updateMethod · 0.45

Tested by

no test coverage detected