FST_code / main_bert_fine_tune.py
jasonfan's picture
2026-03-19
3b2d368 verified
Raw
History Blame Contribute Delete
15 kB
#!/usr/bin/env python3
# main.py - 支持原有流程 + Bert HF fine-tune 模式
import os
import re
import shutil
import csv
import sys
import json
from pathlib import Path
import torch
import torch.distributed as dist
import math
from dotenv import load_dotenv
import hydra
from omegaconf import OmegaConf
# project imports (adjust paths if needed)
from lmr.config import initialize_config
from lmr.tokenizer import Tokenizer
from lmr.models import get_model
from lmr.checkpointing import Checkpointing
from lmr.utils.seed import set_seed
from lmr.training import Bert_Trainer, Trainer
from lmr.generation import Generator
from lmr.benchmark import Benchmark
from lmr.ddp import unwrap_model
# import the BertFineTuneTrainer (ensure this module exists at this path)
# 如果你把 BertFineTuneTrainer 放在 lmr/training/bert_finetune_trainer.py,则如下导入:
try:
from lmr.training.bert_finetune_trainer import BertFineTuneTrainer
except Exception as e:
# 如果没有该文件,提醒并继续(后续会报错)
BertFineTuneTrainer = None
print("⚠️ Warning: Could not import BertFineTuneTrainer: ", e)
DATASET_DIR = Path("datasets")
CHECKPOINT_DIR = Path("/work/jf381/checkpoints")
BENCHMARK_DIR = Path("output")
# -------------------------
# Helper functions (unchanged)
# -------------------------
def load_weight_data(path_obj, device="cpu"):
from safetensors.torch import load_file
model_state = {}
path_obj = Path(path_obj)
if path_obj.is_dir():
index_file = path_obj / "model.safetensors.index.json"
if index_file.exists():
print(f"🔹 Detected sharded safetensors folder: {path_obj.name}")
with open(index_file, 'r') as f:
index_data = json.load(f)
weight_map = index_data.get("weight_map", {})
shards = set(weight_map.values())
for shard_name in shards:
shard_path = path_obj / shard_name
model_state.update(load_file(str(shard_path), device=str(device)))
return model_state
else:
possible = list(path_obj.glob("*.safetensors")) + list(path_obj.glob("*.pt"))
if not possible:
return None
path_obj = possible[0]
if path_obj.suffix == ".safetensors":
print(f"🔹 Loading single safetensors: {path_obj.name}")
return load_file(str(path_obj), device=str(device))
else:
print(f"🔹 Loading pickle (.pt): {path_obj.name}")
ckpt = torch.load(path_obj, map_location=device)
return ckpt.get("model", ckpt.get("state_dict", ckpt))
def load_state_dict_robust(model, checkpoint_path, strict=False):
path_obj = Path(checkpoint_path)
if not path_obj.exists():
print(f"❌ Path not found: {checkpoint_path}")
return False
try:
model_state = load_weight_data(path_obj)
if model_state is None: return False
ckpt_keys_map = {}
for k in model_state.keys():
clean_k = k.replace("module.", "").replace("_orig_mod.", "").replace("model.", "")
ckpt_keys_map[clean_k] = k
target_model = unwrap_model(model)
target_state = target_model.state_dict()
filtered_state = {}
matched_count = 0
for k_target, v_target in target_state.items():
k_target_clean = k_target.replace("module.", "").replace("_orig_mod.", "").replace("model.", "")
if k_target_clean in ckpt_keys_map:
real_ckpt_key = ckpt_keys_map[k_target_clean]
v_ckpt = model_state[real_ckpt_key]
if v_ckpt.shape == v_target.shape:
filtered_state[k_target] = v_ckpt
matched_count += 1
msg = target_model.load_state_dict(filtered_state, strict=strict)
print(f"✅ Success! Loaded {matched_count} parameters. Status: {msg}")
return True
except Exception as e:
print(f"❌ Failed to load checkpoint: {e}")
import traceback
traceback.print_exc()
return False
def setup_model_and_tokenizer(config):
print(f"🔧 Initializing Tokenizer: {config.tokenizer_base}")
tokenizer = Tokenizer(config.tokenizer_base)
print(f"🔧 Initializing Model: {config.model}")
model = get_model(config.model, tokenizer.vocab_size, tokenizer=tokenizer)
return tokenizer, model
def run_generation_task(config, model, tokenizer, output_dir, ckpt_name_tag=""):
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
model.to(device)
model.eval()
generator = Generator(config, model, tokenizer, device=device, output_dir=output_dir)
results = generator.generate()
metrics = {}
if results and "verification" in results:
metrics = {
'acc': results['verification']['accuracy'],
'correct': results['verification']['correct'],
'total': results['verification']['total']
}
if results.get('token_accuracies'):
import numpy as np
metrics['token_acc'] = float(np.mean(results['token_accuracies']))
if ckpt_name_tag:
for fname in ["gsm8k_metrics.txt", "gsm8k_generations.txt"]:
src = output_dir / fname
if src.exists():
dst = output_dir / f"{src.stem}_{ckpt_name_tag}{src.suffix}"
shutil.move(src, dst)
return metrics
# -------------------------
# Train / Generate orchestration
# -------------------------
def generate(config):
tokenizer, model = setup_model_and_tokenizer(config)
ckpt_dir = CHECKPOINT_DIR / config.checkpoint_name
load_mode = getattr(config.benchmark, "checkpoint_mode", "recent")
print(f"🔍 Looking for [{load_mode}] weights in {ckpt_dir}...")
all_items = list(ckpt_dir.iterdir()) if ckpt_dir.exists() else []
checkpoints = [f for f in all_items if ("optim" not in f.name and "sched" not in f.name and f.name != "metrics_summary.csv")]
def sort_key(f):
nums = re.findall(r'\d+', f.name)
return int(nums[-1]) if nums else 0
checkpoints.sort(key=sort_key)
target_ckpt = None
if load_mode == "best":
best_candidates = [f for f in checkpoints if "best" in f.name.lower()]
target_ckpt = best_candidates[0] if best_candidates else (checkpoints[-1] if checkpoints else None)
else:
target_ckpt = checkpoints[-1] if checkpoints else None
if target_ckpt:
print(f"🚀 Found target: {target_ckpt.name}")
success = load_state_dict_robust(model, target_ckpt)
if not success:
print("⚠️ Load failed, check file integrity.")
else:
print(f"❌ No checkpoints found in {ckpt_dir}")
run_generation_task(config, model, tokenizer, output_dir=ckpt_dir)
def generate_all(config):
tokenizer, model = setup_model_and_tokenizer(config)
ckpt_dir = CHECKPOINT_DIR / config.checkpoint_name
all_items = list(ckpt_dir.iterdir()) if ckpt_dir.exists() else []
checkpoints = [f for f in all_items if ("optim" not in f.name and "sched" not in f.name and f.name != "metrics_summary.csv")]
def sort_key(f):
nums = re.findall(r'\d+', f.name)
return int(nums[-1]) if nums else 0
checkpoints.sort(key=sort_key)
print(f"\n🔎 Found {len(checkpoints)} checkpoints.")
summary_path = ckpt_dir / "metrics_summary.csv"
with open(summary_path, mode='w', newline='') as f:
writer = csv.writer(f)
writer.writerow(["checkpoint", "accuracy", "token_accuracy", "correct", "total"])
for ckpt_path in checkpoints:
print(f"\n{'-'*40}\nProcessing: {ckpt_path.name}\n{'-'*40}")
if load_state_dict_robust(model, ckpt_path):
metrics = run_generation_task(config, model, tokenizer, ckpt_dir, ckpt_name_tag=ckpt_path.name)
if metrics:
with open(summary_path, mode='a', newline='') as f:
writer = csv.writer(f)
writer.writerow([
ckpt_path.name,
f"{metrics.get('acc', 0):.4f}",
f"{metrics.get('token_acc', 0):.4f}",
metrics.get('correct', 0),
metrics.get('total', 0)
])
def train_model(config):
# if user wants to run HF Bert fine-tune mode, handle specially
# 配置约定: config.mode == "finetune_bert" 或 config.finetune.bert.enabled == true
use_bert_finetune_mode = False
if getattr(config, "mode", "") == "finetune_bert":
use_bert_finetune_mode = True
elif hasattr(config, "finetune") and getattr(config.finetune, "bert", None) and getattr(config.finetune.bert, "enabled", False):
use_bert_finetune_mode = True
if use_bert_finetune_mode:
if BertFineTuneTrainer is None:
raise RuntimeError("BertFineTuneTrainer not available. 请确认 lmr.training.bert_finetune_trainer.py 存在并且可导入。")
# 构造 BertFineTuneTrainer 需要的 cfg 字典
# 优先使用 config.finetune.bert 下的参数,其次从 config.training/全局取默认
ft_cfg = {}
# required fields sample:
# model_name_or_path, dataset, dataset_config_name (optional), task, num_labels, output_dir, batch_size, eval_batch_size, num_epochs, lr, weight_decay, gradient_accumulation_steps, max_length, max_train_samples, max_eval_samples, fp16, use_ddp
fin = getattr(config, "finetune", None)
if fin and getattr(fin, "bert", None):
fin_bert = fin.bert
else:
# backward compat: maybe config.training contains some entries
fin_bert = getattr(config, "training", {})
# map possible fields — 这里尽量宽容
ft_cfg["model_name_or_path"] = getattr(fin_bert, "model_name_or_path", None) or getattr(config, "model_name", None) or getattr(config, "model", None)
ft_cfg["dataset"] = getattr(fin_bert, "dataset", None) or getattr(config, "dataset", None)
ft_cfg["dataset_config_name"] = getattr(fin_bert, "dataset_config_name", None)
ft_cfg["task"] = getattr(fin_bert, "task", "sentence_pair")
ft_cfg["num_labels"] = getattr(fin_bert, "num_labels", None)
ft_cfg["output_dir"] = str(Path(getattr(fin_bert, "output_dir", config.get("output_dir", "./outputs/finetune")))) if isinstance(config, dict) else str(Path(getattr(fin_bert, "output_dir", "./outputs/finetune")))
ft_cfg["batch_size"] = getattr(fin_bert, "batch_size", getattr(config, "batch_size", 16))
ft_cfg["eval_batch_size"] = getattr(fin_bert, "eval_batch_size", max(32, int(ft_cfg["batch_size"])))
ft_cfg["num_epochs"] = getattr(fin_bert, "num_epochs", getattr(config, "num_epochs", 3))
ft_cfg["lr"] = getattr(fin_bert, "lr", getattr(config, "learning_rate", 2e-5))
ft_cfg["weight_decay"] = getattr(fin_bert, "weight_decay", 0.01)
ft_cfg["gradient_accumulation_steps"] = getattr(fin_bert, "gradient_accumulation_steps", 1)
ft_cfg["max_length"] = getattr(fin_bert, "max_length", 128)
ft_cfg["fp16"] = getattr(fin_bert, "fp16", True)
ft_cfg["use_ddp"] = bool(getattr(config, "distributed", False) or getattr(config, "use_ddp", False))
ft_cfg["seed"] = getattr(config, "seed", 42)
ft_cfg["logging_steps"] = getattr(fin_bert, "logging_steps", 100)
ft_cfg["eval_steps"] = getattr(fin_bert, "eval_steps", 500)
ft_cfg["save_steps"] = getattr(fin_bert, "save_steps", 1000)
ft_cfg["warmup_steps"] = getattr(fin_bert, "warmup_steps", 0)
ft_cfg["max_train_samples"] = getattr(fin_bert, "max_train_samples", None)
ft_cfg["max_eval_samples"] = getattr(fin_bert, "max_eval_samples", None)
ft_cfg["nsp_negatives_ratio"] = getattr(fin_bert, "nsp_negatives_ratio", 1)
# Validation of required args
if not ft_cfg["model_name_or_path"] or not ft_cfg["dataset"]:
raise ValueError("finetune_bert 模式需要在配置中指定 model_name_or_path 和 dataset(Hugging Face dataset id). 示例: finetune.bert.model_name_or_path='bert-base-uncased', finetune.bert.dataset='glue/mrpc'")
# Convert OmegaConf objects to plain types if necessary
if not isinstance(ft_cfg["dataset"], str) and hasattr(ft_cfg["dataset"], "__str__"):
ft_cfg["dataset"] = str(ft_cfg["dataset"])
# Create and run the BertFineTuneTrainer
print(f"🔧 Starting HuggingFace BERT finetune with cfg: model={ft_cfg['model_name_or_path']}, dataset={ft_cfg['dataset']}, task={ft_cfg.get('task')}")
trainer = BertFineTuneTrainer(ft_cfg)
trainer.train()
return
# -------------------------
# 原有的训练流程(不做改动):
# -------------------------
if torch.cuda.is_available() and torch.cuda.device_count() > 1:
if not dist.is_initialized():
dist.init_process_group(backend="nccl")
torch.cuda.set_device(dist.get_rank() % torch.cuda.device_count())
tokenizer, model = setup_model_and_tokenizer(config)
tokenized_dataset_dir = DATASET_DIR / config.tokenizer_base
splits = None
try:
# 尝试调用老的 get_dataset_splits(如果存在)
from lmr.data import get_dataset_splits
splits = get_dataset_splits(config.dataset, 1024, tokenized_dataset_dir)
except Exception:
print("⚠️ get_dataset_splits not available or failed — continuing without it.")
checkpointing = Checkpointing(model, CHECKPOINT_DIR / config.checkpoint_name)
if "bert" in str(config.model).lower():
from lmr.training import Bert_Trainer
trainer = Bert_Trainer(config.training, model, tokenizer, splits, checkpointing, None)
else:
trainer = Trainer(config.training, model, tokenizer, splits, checkpointing)
trainer.train()
# -------------------------
# Entrypoint
# -------------------------
@hydra.main(config_path="config", config_name="config", version_base="1.3")
def main(config):
# config is an OmegaConf object
load_dotenv()
set_seed(config)
initialize_config(config)
mode = getattr(config, "mode", "train")
if mode == "train":
train_model(config)
elif mode == "generate":
generate(config)
elif mode == "generate_all":
generate_all(config)
elif mode == "benchmark":
tokenizer, model = setup_model_and_tokenizer(config)
checkpointing = Checkpointing(model, CHECKPOINT_DIR / config.checkpoint_name)
benchmarking = Benchmark(config.benchmark, model, tokenizer, checkpointing, BENCHMARK_DIR / config.checkpoint_name)
benchmarking.run_benchmarks()
elif mode == "finetune_bert":
# 进入我们上面实现的 fine-tune 分支
train_model(config)
else:
print(f"❌ Unknown mode: {mode}")
if __name__ == "__main__":
main()