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

Method _learner_tree

axlearn/common/learner.py:434–452  ·  view source on GitHub ↗

Returns a tree of the same structure as params where each leaf is the name of the sublearner to apply.

(self, params: Nested[Any])

Source from the content-addressed store, hash-verified

432 def __init__(self, cfg: Config, *, parent: Module):
433 super().__init__(cfg, parent=parent)
434 cfg = self.config
435
436 for name, learner_cfg in cfg.learners.items():
437 # Sub learner should not hold ema.
438 if learner_cfg.ema.decay is not None:
439 raise ValueError(f"Sublearner {name} ema decay is not None.")
440 if name == "ema":
441 raise ValueError("Sublearner name cannot be ema.")
442
443 sub_learner = learner_cfg.set(name=name)
444 self._add_child(name, sub_learner)
445
446 # Check that learners in the rules exist.
447 for _, rule_name in cfg.rules:
448 if rule_name not in cfg.learners:
449 raise ValueError(f"{rule_name} is not found in the known learners.")
450 if cfg.ema.decay is not None:
451 # Create a global model ema.
452 self.ema: PartitionedGradientTransformation = cfg.ema.instantiate()
453
454 def _learner_tree(self, params: Nested[Any]) -> Nested[str]:
455 """Returns a tree of the same structure as params where each leaf is the name of the

Callers 4

initMethod · 0.95
should_applyMethod · 0.95

Calls 3

match_regex_rulesFunction · 0.90
tree_pathsFunction · 0.90
mapMethod · 0.80

Tested by

no test coverage detected