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

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