From 978c88a4edef0035aa9cb05944a0df82817c1cb3 Mon Sep 17 00:00:00 2001 From: Jiarui Li Date: Wed, 22 Jul 2026 14:44:31 +0800 Subject: [PATCH] Log and save model parameter counts --- TrajMixer_设计方案.md | 2 +- test_traj_mixer.py | 11 +++++++++++ train_all_future.py | 13 ++++++++++++- train_next_step.py | 13 ++++++++++++- train_util.py | 14 ++++++++++++++ 5 files changed, 50 insertions(+), 3 deletions(-) diff --git a/TrajMixer_设计方案.md b/TrajMixer_设计方案.md index eb3d10b..48ba7f2 100644 --- a/TrajMixer_设计方案.md +++ b/TrajMixer_设计方案.md @@ -319,7 +319,7 @@ output_init_std: 0.001 group_wise_layer_norm: false ``` -训练时必须将 `model_architecture: traj_mixer_v1` 写入 `train_config.json`。本分支的评估和导出入口只接受带有该标识、且 checkpoint 中包含 TrajMixer 参数张量的模型;其他分支生成的模型应直接拒绝加载。 +训练时必须将 `model_architecture: traj_mixer_v1`、`model_parameter_count` 和 `trainable_parameter_count` 写入 `train_config.json`,并在训练日志中显式打印总参数量与可训练参数量。本分支的评估和导出入口只接受带有该架构标识、且 checkpoint 中包含 TrajMixer 参数张量的模型;其他分支生成的模型应直接拒绝加载。 必须满足: diff --git a/test_traj_mixer.py b/test_traj_mixer.py index c52ad6a..5dc3ae2 100644 --- a/test_traj_mixer.py +++ b/test_traj_mixer.py @@ -8,6 +8,7 @@ from models import ( validate_traj_mixer_config, validate_traj_mixer_state_dict, ) +from train_util import get_model_parameter_counts class TrajMixerTest(unittest.TestCase): @@ -91,6 +92,16 @@ class TrajMixerTest(unittest.TestCase): with self.assertRaisesRegex(ValueError, "divisible"): TrajMixer(n_embd=121, n_head=10, hidden_group=20) + def test_parameter_counts_match_traj_mixer_parameters(self) -> None: + mixer = TrajMixer(n_embd=120, n_head=10, hidden_group=20) + self.assertEqual( + get_model_parameter_counts(mixer), + { + "model_parameter_count": 8_640, + "trainable_parameter_count": 8_640, + }, + ) + if __name__ == "__main__": unittest.main() diff --git a/train_all_future.py b/train_all_future.py index 0fadc99..83879cc 100644 --- a/train_all_future.py +++ b/train_all_future.py @@ -33,6 +33,7 @@ from train_util import ( configure_torch_for_training, create_unique_run_dir, format_extra_info_types, + get_model_parameter_counts, load_extra_info_types_file, resolve_device, save_checkpoint, @@ -437,6 +438,12 @@ def main() -> None: ) model = build_model(args, train_dataset).to(device) + parameter_counts = get_model_parameter_counts(model) + logger.info( + "Model parameters: " + f"total={parameter_counts['model_parameter_count']:,}, " + f"trainable={parameter_counts['trainable_parameter_count']:,}" + ) optimizer = AdamW( model.parameters(), lr=args.base_lr, @@ -446,10 +453,14 @@ def main() -> None: criterion = build_criterion(args, train_dataset) adaptive_lr = args.base_lr * math.sqrt(args.batch_size / 128) + train_metadata = build_metadata( + args, train_dataset, run_name, train_subset, val_subset, test_subset + ) + train_metadata.update(parameter_counts) save_config( args, run_dir / "train_config.json", - extra=build_metadata(args, train_dataset, run_name, train_subset, val_subset, test_subset), + extra=train_metadata, ) best_val = float("inf") diff --git a/train_next_step.py b/train_next_step.py index 22e5f53..23ef9a9 100644 --- a/train_next_step.py +++ b/train_next_step.py @@ -31,6 +31,7 @@ from train_util import ( configure_torch_for_training, create_unique_run_dir, format_extra_info_types, + get_model_parameter_counts, load_extra_info_types_file, resolve_device, save_checkpoint, @@ -599,6 +600,12 @@ def main() -> None: ) model = build_model(args, dataset).to(device) + parameter_counts = get_model_parameter_counts(model) + logger.info( + "Model parameters: " + f"total={parameter_counts['model_parameter_count']:,}, " + f"trainable={parameter_counts['trainable_parameter_count']:,}" + ) readout = build_next_step_readout(args).to(device) criterion = build_next_step_loss(args) optimizer = AdamW( @@ -609,10 +616,14 @@ def main() -> None: ) adaptive_lr = args.base_lr * math.sqrt(args.batch_size / 128) + train_metadata = build_metadata( + args, dataset, run_name, train_subset, val_subset, test_subset + ) + train_metadata.update(parameter_counts) save_config( args, run_dir / "train_config.json", - extra=build_metadata(args, dataset, run_name, train_subset, val_subset, test_subset), + extra=train_metadata, ) best_val = float("inf") diff --git a/train_util.py b/train_util.py index 1fe578c..2a72fcb 100644 --- a/train_util.py +++ b/train_util.py @@ -300,6 +300,20 @@ def build_optimizer(args: Any, model: DeepHealth) -> AdamW: ) +def get_model_parameter_counts(model: torch.nn.Module) -> Dict[str, int]: + """Return stable parameter-count fields for logs and train_config.json.""" + return { + "model_parameter_count": sum( + parameter.numel() for parameter in model.parameters() + ), + "trainable_parameter_count": sum( + parameter.numel() + for parameter in model.parameters() + if parameter.requires_grad + ), + } + + def set_optimizer_lr(optimizer: AdamW, lr: float) -> None: for param_group in optimizer.param_groups: param_group["lr"] = lr