(self, model_params: Nested[OptParam])
| 211 | if cfg.forward_fn_transformation is not None: |
| 212 | self._forward_fn_transformation = maybe_instantiate(cfg.forward_fn_transformation) |
| 213 | else: |
| 214 | self._forward_fn_transformation = lambda fn: fn |
| 215 | |
| 216 | def create_state_partition_specs(self, model_param_specs: Nested[ParameterSpec]) -> Any: |
| 217 | optimizer_model_param_specs = self._get_optimizer_model_params(model_param_specs) |
| 218 | partition_state = dict( |
| 219 | optimizer=self.optimizer.create_state_partition_specs(optimizer_model_param_specs) |
| 220 | ) |
| 221 | if self.config.ema.decay is not None: |
| 222 | partition_state["ema"] = self.ema.partition(model_param_specs) |
| 223 | return partition_state |
| 224 | |
| 225 | def _get_optimizer_model_params(self, model_params: Nested[OptParam]) -> Nested[OptParam]: |
| 226 | should_update_params = self.should_update_with_optimizers(model_params) |
no test coverage detected