From df8906936246c35208376f7b9aa362d639c35aad Mon Sep 17 00:00:00 2001 From: jingyaogong Date: Wed, 7 Jan 2026 23:08:45 +0800 Subject: [PATCH] [update] params log --- trainer/trainer_utils.py | 20 +++++++++----------- 1 file changed, 9 insertions(+), 11 deletions(-) diff --git a/trainer/trainer_utils.py b/trainer/trainer_utils.py index a50737b..58d0549 100644 --- a/trainer/trainer_utils.py +++ b/trainer/trainer_utils.py @@ -16,17 +16,15 @@ from model.model_minimind import MiniMindForCausalLM def get_model_params(model, config): total = sum(p.numel() for p in model.parameters()) / 1e6 - if getattr(config, 'use_moe', False): - n_routed = getattr(config, 'n_routed_experts', getattr(config, 'num_experts', 0)) - n_active = getattr(config, 'num_experts_per_tok', 0) - n_shared = getattr(config, 'n_shared_experts', 0) - expert = sum(p.numel() for n, p in model.named_parameters() if 'mlp.experts.0.' in n) / 1e6 - shared_expert = sum(p.numel() for n, p in model.named_parameters() if 'mlp.shared_experts.0.' in n) / 1e6 - base = total - (expert * n_routed) - (shared_expert * n_shared) - active = base + (expert * n_active) + (shared_expert * n_shared) - Logger(f'Model Params: {total:.2f}M-A{active:.2f}M') - else: - Logger(f'Model Params: {total:.2f}M') + n_routed = getattr(config, 'n_routed_experts', getattr(config, 'num_experts', 0)) + n_active = getattr(config, 'num_experts_per_tok', 0) + n_shared = getattr(config, 'n_shared_experts', 0) + expert = sum(p.numel() for n, p in model.named_parameters() if 'mlp.experts.0.' in n) / 1e6 + shared_expert = sum(p.numel() for n, p in model.named_parameters() if 'mlp.shared_experts.0.' in n) / 1e6 + base = total - (expert * n_routed) - (shared_expert * n_shared) + active = base + (expert * n_active) + (shared_expert * n_shared) + if active < total: Logger(f'Model Params: {total:.2f}M-A{active:.2f}M') + else: Logger(f'Model Params: {total:.2f}M') def is_main_process():