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

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