mirror of
https://github.com/jingyaogong/minimind.git
synced 2026-09-28 05:47:23 +00:00
210 lines
9.2 KiB
Python
210 lines
9.2 KiB
Python
"""
|
||
训练工具函数集合
|
||
"""
|
||
import os
|
||
import sys
|
||
__package__ = "trainer"
|
||
sys.path.append(os.path.abspath(os.path.join(os.path.dirname(__file__), '..')))
|
||
import ast
|
||
import operator
|
||
import random
|
||
import math
|
||
import numpy as np
|
||
import torch
|
||
import torch.distributed as dist
|
||
from torch.nn.parallel import DistributedDataParallel
|
||
from torch.utils.data import Sampler
|
||
from transformers import AutoTokenizer, AutoModel, AutoModelForSequenceClassification
|
||
from model.model_minimind import MiniMindForCausalLM
|
||
|
||
def get_model_params(model, config):
|
||
total = sum(p.numel() for p in model.parameters()) / 1e6
|
||
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():
|
||
return not dist.is_initialized() or dist.get_rank() == 0
|
||
|
||
|
||
def Logger(content):
|
||
if is_main_process():
|
||
print(content)
|
||
|
||
|
||
def get_lr(current_step, total_steps, lr):
|
||
return lr*(0.1 + 0.45*(1 + math.cos(math.pi * current_step / total_steps)))
|
||
|
||
|
||
def init_distributed_mode():
|
||
if int(os.environ.get("RANK", -1)) == -1:
|
||
return 0 # 非DDP模式
|
||
|
||
dist.init_process_group(backend="nccl")
|
||
local_rank = int(os.environ["LOCAL_RANK"])
|
||
torch.cuda.set_device(local_rank)
|
||
return local_rank
|
||
|
||
|
||
def setup_seed(seed: int):
|
||
random.seed(seed)
|
||
np.random.seed(seed)
|
||
torch.manual_seed(seed)
|
||
torch.cuda.manual_seed(seed)
|
||
torch.cuda.manual_seed_all(seed)
|
||
torch.backends.cudnn.deterministic = True
|
||
torch.backends.cudnn.benchmark = False
|
||
|
||
def lm_checkpoint(lm_config, weight='full_sft', model=None, optimizer=None, epoch=0, step=0, wandb=None, save_dir='../checkpoints', **kwargs):
|
||
os.makedirs(save_dir, exist_ok=True)
|
||
moe_path = '_moe' if lm_config.use_moe else ''
|
||
ckp_path = f'{save_dir}/{weight}_{lm_config.hidden_size}{moe_path}.pth'
|
||
resume_path = f'{save_dir}/{weight}_{lm_config.hidden_size}{moe_path}_resume.pth'
|
||
|
||
if model is not None:
|
||
raw_model = model.module if isinstance(model, DistributedDataParallel) else model
|
||
raw_model = getattr(raw_model, '_orig_mod', raw_model)
|
||
state_dict = raw_model.state_dict()
|
||
state_dict = {k: v.half().cpu() for k, v in state_dict.items()}
|
||
ckp_tmp = ckp_path + '.tmp'
|
||
torch.save(state_dict, ckp_tmp)
|
||
os.replace(ckp_tmp, ckp_path)
|
||
wandb_id = None
|
||
if wandb:
|
||
if hasattr(wandb, 'get_run'):
|
||
run = wandb.get_run()
|
||
wandb_id = getattr(run, 'id', None) if run else None
|
||
else:
|
||
wandb_id = getattr(wandb, 'id', None)
|
||
|
||
resume_data = {
|
||
'model': state_dict,
|
||
'optimizer': optimizer.state_dict(),
|
||
'epoch': epoch,
|
||
'step': step,
|
||
'world_size': dist.get_world_size() if dist.is_initialized() else 1,
|
||
'wandb_id': wandb_id
|
||
}
|
||
for key, value in kwargs.items():
|
||
if value is not None:
|
||
if hasattr(value, 'state_dict'):
|
||
raw_value = value.module if isinstance(value, DistributedDataParallel) else value
|
||
raw_value = getattr(raw_value, '_orig_mod', raw_value)
|
||
resume_data[key] = raw_value.state_dict()
|
||
else:
|
||
resume_data[key] = value
|
||
|
||
resume_tmp = resume_path + '.tmp'
|
||
torch.save(resume_data, resume_tmp)
|
||
os.replace(resume_tmp, resume_path)
|
||
del state_dict, resume_data
|
||
torch.cuda.empty_cache()
|
||
else: # 加载模式
|
||
if os.path.exists(resume_path):
|
||
ckp_data = torch.load(resume_path, map_location='cpu')
|
||
saved_ws = ckp_data.get('world_size', 1)
|
||
current_ws = dist.get_world_size() if dist.is_initialized() else 1
|
||
if saved_ws != current_ws:
|
||
ckp_data['step'] = ckp_data['step'] * saved_ws // current_ws
|
||
Logger(f'GPU数量变化({saved_ws}→{current_ws}),step已自动转换为{ckp_data["step"]}')
|
||
return ckp_data
|
||
return None
|
||
|
||
|
||
def init_model(lm_config, from_weight='pretrain', tokenizer_path='../model', save_dir='../out', device='cuda'):
|
||
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path)
|
||
model = MiniMindForCausalLM(lm_config)
|
||
|
||
if from_weight!= 'none':
|
||
moe_suffix = '_moe' if lm_config.use_moe else ''
|
||
weight_path = f'{save_dir}/{from_weight}_{lm_config.hidden_size}{moe_suffix}.pth'
|
||
weights = torch.load(weight_path, map_location=device)
|
||
model.load_state_dict(weights, strict=False)
|
||
|
||
get_model_params(model, lm_config)
|
||
Logger(f'Trainable Params: {sum(p.numel() for p in model.parameters() if p.requires_grad) / 1e6:.3f}M')
|
||
return model.to(device), tokenizer
|
||
|
||
|
||
class SkipBatchSampler(Sampler):
|
||
def __init__(self, sampler, batch_size, skip_batches=0):
|
||
self.sampler = sampler
|
||
self.batch_size = batch_size
|
||
self.skip_batches = skip_batches
|
||
|
||
def __iter__(self):
|
||
batch = []
|
||
skipped = 0
|
||
for idx in self.sampler:
|
||
batch.append(idx)
|
||
if len(batch) == self.batch_size:
|
||
if skipped < self.skip_batches:
|
||
skipped += 1
|
||
batch = []
|
||
continue
|
||
yield batch
|
||
batch = []
|
||
if len(batch) > 0 and skipped >= self.skip_batches:
|
||
yield batch
|
||
|
||
def __len__(self):
|
||
total_batches = (len(self.sampler) + self.batch_size - 1) // self.batch_size
|
||
return max(0, total_batches - self.skip_batches)
|
||
|
||
|
||
class LMForRewardModel:
|
||
def __init__(self, model_path, device="cuda", dtype=torch.float16):
|
||
self.tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
|
||
self.model = AutoModel.from_pretrained(model_path, torch_dtype=dtype, trust_remote_code=True)
|
||
self.model = self.model.to(device).eval()
|
||
self.device = device
|
||
|
||
@torch.no_grad()
|
||
def get_score(self, messages, response):
|
||
history_text = "\n".join([f"{m['role']}: {m['content']}" for m in messages[:-1]])
|
||
last_query = messages[-1]['content'] if messages else ""
|
||
message_context = f"{history_text}\n以上是对话历史。我的新问题是:\n{last_query}" if history_text else last_query
|
||
eval_messages = [
|
||
{"role": "user", "content": message_context},
|
||
{"role": "assistant", "content": response}
|
||
]
|
||
score = self.model.get_score(self.tokenizer, eval_messages)
|
||
return max(min(score, 3.0), -3.0)
|
||
|
||
|
||
# ===== 数学表达式安全求值:替代 eval,只放行算术运算与 math 白名单(长度上限 512) =====
|
||
def safe_math_eval(expression):
|
||
"""对模型生成的数学表达式求值:支持 + - * / // % **、math 白名单函数与 pi/e/tau 常量。"""
|
||
def pow_guard(base, exp): # 幂运算统一走这里,拦截 9**9**9 / 10**99999 这类把进程算死的输入
|
||
if abs(exp) > 1e4 or (abs(base) > 1 and abs(exp) * math.log10(abs(base)) > 100): raise ValueError('幂运算结果过大')
|
||
return base ** exp
|
||
def resolve(node): # 同时兼容 sqrt(4) 与 math.sqrt(4) 两种写法
|
||
if isinstance(node, ast.Name): return node.id
|
||
if isinstance(node, ast.Attribute) and getattr(node.value, 'id', '') == 'math': return node.attr
|
||
def walk(node):
|
||
if isinstance(node, ast.Constant) and type(node.value) in (int, float): return node.value
|
||
name = resolve(node)
|
||
if name in consts: return consts[name]
|
||
if isinstance(node, ast.UnaryOp): return unary_ops[type(node.op)](walk(node.operand))
|
||
if isinstance(node, ast.BinOp): return bin_ops[type(node.op)](walk(node.left), walk(node.right))
|
||
if isinstance(node, ast.Call) and not node.keywords: return funcs[resolve(node.func)](*map(walk, node.args))
|
||
raise ValueError(f'不支持的表达式语法: {type(node).__name__}')
|
||
consts = {'pi': math.pi, 'e': math.e, 'tau': math.tau}
|
||
funcs = {'pow': pow_guard, **{n: getattr(math, n) for n in 'sqrt exp log log2 log10 sin cos tan asin acos atan atan2 sinh cosh tanh floor ceil trunc fabs fmod hypot gcd degrees radians'.split()}}
|
||
bin_ops = {ast.Add: operator.add, ast.Sub: operator.sub, ast.Mult: operator.mul, ast.Div: operator.truediv, ast.FloorDiv: operator.floordiv, ast.Mod: operator.mod, ast.Pow: pow_guard}
|
||
unary_ops = {ast.UAdd: operator.pos, ast.USub: operator.neg}
|
||
chars = str.maketrans({'^': '**', '×': '*', '÷': '/', '−': '-', '²': '**2', '³': '**3', '(': '(', ')': ')'})
|
||
expr = str(expression).translate(chars).strip()
|
||
if not expr or len(expr) > 512: raise ValueError('表达式为空或过长')
|
||
try:
|
||
return walk(ast.parse(expr, mode='eval').body)
|
||
except (KeyError, SyntaxError, TypeError):
|
||
raise ValueError('不支持的表达式语法') from None
|