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)
| 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( |
no test coverage detected