Tests composite learner with two sub learners for weight/bias respectively.
(self)
| 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")) |
nothing calls this directly
no test coverage detected