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

Method test_composite_learner

axlearn/common/trainer_test.py:1050–1093  ·  view source on GitHub ↗

Tests composite learner with two sub learners for weight/bias respectively.

(self)

Source from the content-addressed store, hash-verified

1048 model_cfg = DummyModel.default_config().set(dtype=jnp.float32)
1049
1050 def checkpoint_if_all_evalers_run(evaler_names: Sequence[str]) -> CheckpointPolicy:
1051 def fn(*, step: int, evaler_summaries: dict[str, Any]):
1052 del step
1053 for evaler_name in evaler_names:
1054 if evaler_summaries.get(evaler_name) is None:
1055 return False
1056 return True
1057
1058 return fn
1059
1060 cfg: SpmdTrainer.Config = SpmdTrainer.default_config().set(
1061 name="test_trainer",
1062 dir=tempfile.mkdtemp(),
1063 mesh_axis_names=("data", "model"),
1064 mesh_shape=(1, 1),
1065 model=model_cfg,
1066 input=DummyInput.default_config(),
1067 learner=learner.Learner.default_config().set(
1068 optimizer=config_for_function(optimizers.sgd_optimizer).set(
1069 learning_rate=0.1, decouple_weight_decay=True
1070 ),
1071 ),
1072 max_step=8,
1073 checkpointer=Checkpointer.default_config().set(
1074 save_policy=config_for_function(checkpoint_if_all_evalers_run).set(
1075 evaler_names=["eval_every_2", "eval_every_3"]
1076 ),
1077 ),
1078 evalers=dict(
1079 eval_every_2=SpmdEvaler.default_config().set(
1080 input=DummyInput.default_config().set(total_num_batches=1),
1081 eval_policy=config_for_function(eval_every_n_steps_policy).set(n=2),
1082 ),
1083 eval_every_3=SpmdEvaler.default_config().set(
1084 input=DummyInput.default_config().set(total_num_batches=2),
1085 eval_policy=config_for_function(eval_every_n_steps_policy).set(n=3),
1086 ),
1087 ),
1088 save_input_iterator=save_input_iterator,
1089 )
1090 cfg.checkpointer.storage.max_concurrent_gb = max_concurrent_gb
1091
1092 # Run trainer.
1093 trainer: SpmdTrainer = cfg.instantiate(parent=None)
1094 trainer.run(prng_key=jax.random.PRNGKey(123))
1095
1096 assert os.path.exists(os.path.join(cfg.dir, "trainer_state_tree.txt"))

Callers

nothing calls this directly

Calls 8

config_for_functionFunction · 0.90
flatten_itemsFunction · 0.90
setMethod · 0.45
default_configMethod · 0.45
instantiateMethod · 0.45
meshMethod · 0.45
initMethod · 0.45
runMethod · 0.45

Tested by

no test coverage detected