Log and save model parameter counts

This commit is contained in:
2026-07-22 14:44:31 +08:00
parent db0947ce9d
commit 978c88a4ed
5 changed files with 50 additions and 3 deletions

View File

@@ -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 参数张量的模型;其他分支生成的模型应直接拒绝加载。
必须满足:

View File

@@ -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()

View File

@@ -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")

View File

@@ -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")

View File

@@ -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