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
|
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_config,
|
||||||
validate_traj_mixer_state_dict,
|
validate_traj_mixer_state_dict,
|
||||||
)
|
)
|
||||||
|
from train_util import get_model_parameter_counts
|
||||||
|
|
||||||
|
|
||||||
class TrajMixerTest(unittest.TestCase):
|
class TrajMixerTest(unittest.TestCase):
|
||||||
@@ -91,6 +92,16 @@ class TrajMixerTest(unittest.TestCase):
|
|||||||
with self.assertRaisesRegex(ValueError, "divisible"):
|
with self.assertRaisesRegex(ValueError, "divisible"):
|
||||||
TrajMixer(n_embd=121, n_head=10, hidden_group=20)
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -33,6 +33,7 @@ from train_util import (
|
|||||||
configure_torch_for_training,
|
configure_torch_for_training,
|
||||||
create_unique_run_dir,
|
create_unique_run_dir,
|
||||||
format_extra_info_types,
|
format_extra_info_types,
|
||||||
|
get_model_parameter_counts,
|
||||||
load_extra_info_types_file,
|
load_extra_info_types_file,
|
||||||
resolve_device,
|
resolve_device,
|
||||||
save_checkpoint,
|
save_checkpoint,
|
||||||
@@ -437,6 +438,12 @@ def main() -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
model = build_model(args, train_dataset).to(device)
|
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(
|
optimizer = AdamW(
|
||||||
model.parameters(),
|
model.parameters(),
|
||||||
lr=args.base_lr,
|
lr=args.base_lr,
|
||||||
@@ -446,10 +453,14 @@ def main() -> None:
|
|||||||
criterion = build_criterion(args, train_dataset)
|
criterion = build_criterion(args, train_dataset)
|
||||||
adaptive_lr = args.base_lr * math.sqrt(args.batch_size / 128)
|
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(
|
save_config(
|
||||||
args,
|
args,
|
||||||
run_dir / "train_config.json",
|
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")
|
best_val = float("inf")
|
||||||
|
|||||||
@@ -31,6 +31,7 @@ from train_util import (
|
|||||||
configure_torch_for_training,
|
configure_torch_for_training,
|
||||||
create_unique_run_dir,
|
create_unique_run_dir,
|
||||||
format_extra_info_types,
|
format_extra_info_types,
|
||||||
|
get_model_parameter_counts,
|
||||||
load_extra_info_types_file,
|
load_extra_info_types_file,
|
||||||
resolve_device,
|
resolve_device,
|
||||||
save_checkpoint,
|
save_checkpoint,
|
||||||
@@ -599,6 +600,12 @@ def main() -> None:
|
|||||||
)
|
)
|
||||||
|
|
||||||
model = build_model(args, dataset).to(device)
|
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)
|
readout = build_next_step_readout(args).to(device)
|
||||||
criterion = build_next_step_loss(args)
|
criterion = build_next_step_loss(args)
|
||||||
optimizer = AdamW(
|
optimizer = AdamW(
|
||||||
@@ -609,10 +616,14 @@ def main() -> None:
|
|||||||
)
|
)
|
||||||
adaptive_lr = args.base_lr * math.sqrt(args.batch_size / 128)
|
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(
|
save_config(
|
||||||
args,
|
args,
|
||||||
run_dir / "train_config.json",
|
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")
|
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:
|
def set_optimizer_lr(optimizer: AdamW, lr: float) -> None:
|
||||||
for param_group in optimizer.param_groups:
|
for param_group in optimizer.param_groups:
|
||||||
param_group["lr"] = lr
|
param_group["lr"] = lr
|
||||||
|
|||||||
Reference in New Issue
Block a user