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])
| 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 |