Log and save model parameter counts
This commit is contained in:
@@ -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 参数张量的模型;其他分支生成的模型应直接拒绝加载。
|
||||
|
||||
必须满足:
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user