File size: 11,564 Bytes
3b2d368
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
import os
import re
import shutil
import csv
import json
from pathlib import Path

import torch
import torch.distributed as dist
from dotenv import load_dotenv
import hydra
from omegaconf import OmegaConf

# 项目内部引用 - 保持你原有的导入结构
from lmr.config import initialize_config
from lmr.tokenizer import Tokenizer
from lmr.models import get_model
from lmr.data import get_dataset_splits
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

DATASET_DIR = Path("datasets")
CHECKPOINT_DIR = Path("/work/jf381/checkpoints")
BENCHMARK_DIR = Path("output")

# =============================================================================
# 1. 强力权重加载逻辑:处理 pytorch_model.bin
# =============================================================================

def load_state_dict_robust(model, checkpoint_path, strict=False):
    """
    自动处理任意 *.index.json 的 HF 分片权重 / 单文件 bin / safetensors
    """
    path_obj = Path(checkpoint_path)
    state_dict = {}

    try:
        # ============================================================
        # 1️⃣ 自动查找 *.index.json(而不是写死文件名)
        # ============================================================
        index_files = list(path_obj.glob("*.index.json")) if path_obj.is_dir() else []

        if index_files:
            index_file = index_files[0]  # 默认取第一个
            print(f"🧩 检测到分片权重索引: {index_file.name}")

            with open(index_file, "r") as f:
                index_data = json.load(f)

            weight_files = set(index_data["weight_map"].values())
            for wf in sorted(weight_files):
                wf_path = path_obj / wf
                if not wf_path.exists():
                    raise FileNotFoundError(f"Missing shard file: {wf_path}")
                print(f"  📦 加载分片: {wf}")
                part_dict = torch.load(wf_path, map_location="cpu")
                state_dict.update(part_dict)

        else:
            # ============================================================
            # 2️⃣ 单文件权重逻辑(bin / pt / safetensors)
            # ============================================================
            if path_obj.is_dir():
                candidates = (
                    list(path_obj.glob("pytorch_model.bin")) +
                    list(path_obj.glob("*.safetensors")) +
                    list(path_obj.glob("*.pt"))
                )
                weight_file = candidates[0] if candidates else None
            else:
                weight_file = path_obj

            if not weight_file or not weight_file.exists():
                print(f"❌ 找不到权重文件于: {checkpoint_path}")
                return False

            print(f"🔹 正在加载单文件权重: {weight_file.name}")
            if weight_file.suffix == ".safetensors":
                from safetensors.torch import load_file
                state_dict = load_file(str(weight_file), device="cpu")
            else:
                state_dict = torch.load(weight_file, map_location="cpu")

        # ============================================================
        # 3️⃣ 通用清洗 / key 对齐 / shape 校验(保持你原逻辑)
        # ============================================================
        if isinstance(state_dict, dict) and ("model" in state_dict or "state_dict" in state_dict):
            state_dict = state_dict.get("model", state_dict.get("state_dict"))

        target_model = unwrap_model(model)
        target_state = target_model.state_dict()

        clean_state = {}
        for k, v in state_dict.items():
            name = (
                k.replace("module.", "")
                 .replace("_orig_mod.", "")
                 .replace("model.", "")
            )
            clean_state[name] = v

        filtered_state = {}
        matched_count = 0
        for k_target, v_target in target_state.items():
            k_clean = (
                k_target.replace("module.", "")
                        .replace("_orig_mod.", "")
                        .replace("model.", "")
            )
            if k_clean in clean_state:
                v_ckpt = clean_state[k_clean]
                if v_ckpt.shape == v_target.shape:
                    filtered_state[k_target] = v_ckpt
                    matched_count += 1
                else:
                    print(f"⚠️ 形状不匹配跳过: {k_target} ({v_target.shape} vs {v_ckpt.shape})")

        msg = target_model.load_state_dict(filtered_state, strict=strict)

        if hasattr(target_model, "tie_weights"):
            target_model.tie_weights()
            print("🔗 Weights tied successfully.")

        print(f"✅ 成功加载 {matched_count} 个参数 | 状态: {msg}")
        return True

    except Exception as e:
        print(f"❌ 加载失败: {e}")
        import traceback
        traceback.print_exc()
        return False


# =============================================================================
# 2. 模型初始化逻辑
# =============================================================================

def setup_model_and_tokenizer(config):
    """
    使用项目本地代码初始化模型架构
    """
    print(f"🔧 正在初始化 Tokenizer: {config.tokenizer_base}")
    tokenizer = Tokenizer(config.tokenizer_base)
    
    # 检查本地是否有最新的 config.json,如果有,可以根据它调整 vocab_size
    ckpt_dir = CHECKPOINT_DIR / config.checkpoint_name
    local_config_path = next(ckpt_dir.glob("**/config.json"), None)
    
    vocab_size = tokenizer.vocab_size
    if local_config_path:
        with open(local_config_path, 'r') as f:
            local_meta = json.load(f)
            if "vocab_size" in local_meta:
                vocab_size = local_meta["vocab_size"]
                print(f"📖 从本地 Config 读入 Vocab Size: {vocab_size}")

    print(f"🔧 正在初始化模型架构: {config.model}")
    model = get_model(config.model, vocab_size, tokenizer=tokenizer)
    return tokenizer, model

# =============================================================================
# 3. 运行模式逻辑 (Train / Generate)
# =============================================================================

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()
    
    # 重命名结果文件以区分不同 checkpoint
    if ckpt_name_tag and results:
        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 results

def generate(config):
    tokenizer, model = setup_model_and_tokenizer(config)
    ckpt_dir = CHECKPOINT_DIR / config.checkpoint_name
    
    # 1. 搜集目录下所有有效的文件夹
    all_items = [f for f in ckpt_dir.iterdir() if f.is_dir()]
    # 过滤掉不含 epoch 的杂质文件夹
    checkpoints = [f for f in all_items if "epoch" in f.name]
    
    # 2. 根据模式选取
    load_mode = getattr(config.benchmark, "checkpoint_mode", "recent")
    
    if load_mode == "best":
        # 筛选名字里带 'best' 的
        best_ckpts = [f for f in checkpoints if "best" in f.name]
        if best_ckpts:
            # 如果有多个 best,选 epoch 最大的那个
            best_ckpts.sort(key=lambda f: int(re.findall(r'\d+', f.name)[-1]) if re.findall(r'\d+', f.name) else 0)
            target_ckpt = best_ckpts[-1]
        else:
            print("⚠️ No 'best' checkpoint found, falling back to most recent.")
            target_ckpt = sorted(checkpoints)[-1] if checkpoints else None
    else:
        # 默认选最近的 (recent)
        checkpoints.sort(key=lambda f: int(re.findall(r'\d+', f.name)[-1]) if re.findall(r'\d+', f.name) else 0)
        target_ckpt = checkpoints[-1] if checkpoints else None

    if target_ckpt:
        print(f"🚀 Target checkpoint identified: {target_ckpt.name}")
        load_state_dict_robust(model, target_ckpt)
    else:
        print(f"❌ No valid checkpoints found in {ckpt_dir}")
        return

    run_generation_task(config, model, tokenizer, output_dir=ckpt_dir)
    
def generate_all(config):
    """扫描所有 checkpoint 并批量生成"""
    tokenizer, model = setup_model_and_tokenizer(config)
    ckpt_dir = CHECKPOINT_DIR / config.checkpoint_name
    
    all_ckpts = [d for d in ckpt_dir.iterdir() if d.is_dir() and "epoch" in d.name]
    all_ckpts.sort(key=lambda x: int(re.findall(r'\d+', x.name)[-1]) if re.findall(r'\d+', x.name) else 0)
    
    summary_path = ckpt_dir / "metrics_summary.csv"
    with open(summary_path, mode='w', newline='') as f:
        writer = csv.writer(f)
        writer.writerow(["checkpoint", "status"])

    for ckpt_path in all_ckpts:
        print(f"\n{'-'*50}\n处理中: {ckpt_path.name}\n{'-'*50}")
        if load_state_dict_robust(model, ckpt_path):
            run_generation_task(config, model, tokenizer, ckpt_dir, ckpt_name_tag=ckpt_path.name)
            with open(summary_path, mode='a', newline='') as f:
                csv.writer(f).writerow([ckpt_path.name, "Success"])

def train_model(config):
    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 = get_dataset_splits(config.dataset, 1024, tokenized_dataset_dir)
    checkpointing = Checkpointing(model, CHECKPOINT_DIR / config.checkpoint_name)
    
    trainer_cls = Bert_Trainer if "bert" in str(config.model).lower() else Trainer
    trainer = trainer_cls(config.training, model, tokenizer, splits, checkpointing)
    trainer.train()

# =============================================================================
# 4. 主入口
# =============================================================================

@hydra.main(config_path="config", config_name="config", version_base="1.3")
def main(config):
    load_dotenv()
    set_seed(config)
    initialize_config(config)
    
    # 针对 H200 环境下的调试设置
    os.environ['CUDA_LAUNCH_BLOCKING'] = "1"
    os.environ['TORCH_USE_CUDA_DSA'] = "1"

    mode = config.mode
    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()
    else:
        print(f"❌ 未知模式: {mode}")

if __name__ == "__main__":
    main()