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

Method init

axlearn/common/learner.py:213–223  ·  view source on GitHub ↗
(self, model_params: Nested[OptParam])

Source from the content-addressed store, hash-verified

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)

Callers 1

initMethod · 0.45

Calls 4

_update_typesMethod · 0.95
pathMethod · 0.45

Tested by

no test coverage detected