#!/usr/bin/env python3 """HuggingFace Space 训练节点 — 多进程训练 + 推理 + ckpt 上传. 功能: 1. 启动时: clone 代码, 下载数据 + ckpt, 编译, 后台多进程训练 2. Gradio UI: 显示训练进度 + 推理测试 3. 每 500 步: 停所有 worker → merge 多进程 ckpt → 上传 → 下载对方 Space → merge → 重启 多进程配置 (ZeroGPU 容器 2 vCPU): - N_WORKERS=2 个进程, 每进程 OMP_NUM_THREADS=1 - 每个 worker 独立 LAL_WORKER_ID, 独立 ckpt (model_mp_{id}.ste) - 比 1 进程 2 线程快 ~1.7x (无线程竞争, 独立 L1/L2 cache) """ import os import threading import subprocess import time import gradio as gr import spaces from huggingface_hub import HfApi, hf_hub_download # === 固化配置 === WORK_DIR = "/home/user/lalmodel" TRAIN_LOG = f"{WORK_DIR}/train.log" # worker 0 的日志 (UI 显示用) CKPT_PATH = f"{WORK_DIR}/model_dialogue.ste" # merged ckpt 路径 (上传/续训用) DATA_REPO = "gasschina/laloop" # 数据+源码 repo (原 lalmodel-code; 需 src/* 布局且含 data/cci2_clean.bin16 → 只有 laloop 同时满足, lalmodel 是 models/ 布局无 src/) CKPT_REPO = "gasschina/lalmodel" # ckpt 仓 (原 lal-ckpts, 已收编进 ckpts/) HF_TOKEN = os.environ.get("HF_TOKEN", "") UPLOAD_INTERVAL = 500 # 每 500 步上传 ckpt N_WORKERS = int(os.environ.get("LAL_N_WORKERS", "2")) # ZeroGPU 容器 2 vCPU os.makedirs(WORK_DIR, exist_ok=True) def log(msg): print(f"[SPACE] {msg}", flush=True) # ============================================================ # 多进程训练控制 # ============================================================ def start_workers(): """启动 N_WORKERS 个训练进程.""" # 强制杀所有旧 worker (SIGKILL, 多次确认) os.system("pkill -9 -f ste_train 2>/dev/null") time.sleep(2) os.system("pkill -9 -f ste_train 2>/dev/null") time.sleep(1) # 确认没有残留 import subprocess result = subprocess.run(["pgrep", "-f", "ste_train"], capture_output=True, text=True) if result.stdout.strip(): log(f"WARNING: {len(result.stdout.strip().split())} ste_train still running after pkill") for i in range(N_WORKERS): env = (f"OMP_NUM_THREADS=1 OMP_MAX_ACTIVE_LEVELS=1 LAL_WORKER_ID={i} " f"LAL_N_WORKERS={N_WORKERS}") log_file = f"{WORK_DIR}/train.log" if i == 0 else f"{WORK_DIR}/worker_{i}.log" err_file = f"{WORK_DIR}/worker_{i}.err" # stdbuf -oL: 行缓冲, 避免 C 程序 stdout 4KB 缓冲导致日志看不到 os.system( f"cd {WORK_DIR} && {env} " f"stdbuf -oL nohup ./ste_train > {log_file} 2> {err_file} &" ) log(f"worker {i} started (log: {os.path.basename(log_file)})") time.sleep(0.5) def stop_workers(): os.system("pkill -f ste_train 2>/dev/null") time.sleep(3) def merge_local_workers(): """merge 本 Space 的 N_WORKERS 个 ckpt 到 model_dialogue.ste. step-weighted average, 调用 ste_train --merge. 读 ckpt_mp_{id} (周期 save 的 ckpt, save_every=20), 而非 model_mp_{id}.ste (训练结束才 save). """ ckpts = [] for i in range(N_WORKERS): # 优先用 ckpt_mp_{id} (周期 save, 有数据) p = f"{WORK_DIR}/ckpt_mp_{i}" if os.path.exists(p) and os.path.getsize(p) > 100_000_000: ckpts.append(p) else: # 后备: model_mp_{id}.ste (训练结束才存在) p2 = f"{WORK_DIR}/model_mp_{i}.ste" if os.path.exists(p2) and os.path.getsize(p2) > 100_000_000: ckpts.append(p2) if len(ckpts) < 1: log("local merge: no worker ckpts found") return False if len(ckpts) == 1: # 只有一个 worker 有 ckpt, 直接复制 os.system(f"cp {ckpts[0]} {CKPT_PATH}") log(f"local merge: single worker, copied") return True merged = f"{WORK_DIR}/merged_local.ste" args = " ".join(ckpts) os.system(f"{WORK_DIR}/ste_train --merge {merged} {args} 2>&1 | tail -3") if os.path.exists(merged) and os.path.getsize(merged) > 100_000_000: os.system(f"cp {merged} {CKPT_PATH}") os.remove(merged) log(f"local merge: {len(ckpts)} workers merged") return True log("local merge: failed, keeping worker_0 ckpt") if os.path.exists(ckpts[0]): os.system(f"cp {ckpts[0]} {CKPT_PATH}") return False # ============================================================ # 初始化 # ============================================================ def setup(): """下载代码 + 数据 + ckpt, 编译, 启动多进程训练.""" log("Setting up (multi-process mode)...") # 1. download source code from HF (gasschina/laloop src/) — no github dependency # 2026-08-17: github clone was failing inside ZeroGPU container # ("fatal: could not read Username for 'https://github.com'") because the # container doesn't have direct github access. Switched to HF snapshot_download # from gasschina/laloop (which is the same HF we're already on). if not os.path.exists(f"{WORK_DIR}/Makefile"): try: from huggingface_hub import snapshot_download snapshot_download( repo_id=DATA_REPO, repo_type="model", local_dir=f"{WORK_DIR}/_hf_src", allow_patterns=["src/*"], ) # Move src/* contents up to WORK_DIR import shutil src_root = f"{WORK_DIR}/_hf_src/src" for item in os.listdir(src_root): src = os.path.join(src_root, item) dst = os.path.join(WORK_DIR, item) if os.path.exists(dst): shutil.rmtree(dst) if os.path.isdir(dst) else os.remove(dst) shutil.move(src, dst) shutil.rmtree(f"{WORK_DIR}/_hf_src") log(f"source code downloaded from HF (gasschina/laloop/src/)") except Exception as e: log(f"FATAL: source download from HF failed: {e}") return else: log("source already present in WORK_DIR, skipping download") # 2. 下载数据 (v27: CCI2+CNews 14.4亿token, LALT16; 旧wiki已弃用) if not os.path.exists(f"{WORK_DIR}/data/cci2_clean.bin16"): try: p = hf_hub_download(repo_id=DATA_REPO, filename="data/cci2_clean.bin16", repo_type="model") os.system(f"mkdir -p {WORK_DIR}/data && cp {p} {WORK_DIR}/data/") log("cci2_clean.bin16 downloaded (CCI2+CNews, 14.4亿 token, LALT16)") except Exception as e: log(f"data download failed: {e}") # 3. v21: fresh start — 杀旧 worker + 删所有 ckpt + 强制 fresh # 问题: HF Space 持久化存储, 重启后旧 ckpt_mp_* 还在. # 且旧 worker 进程可能没被立即杀掉, 会 save 新 ckpt 覆盖删除. # 修复: 反复 pkill + 反复删除, 确保没有 ckpt 残留. log("Killing any leftover ste_train processes (aggressive)...") for _ in range(3): os.system("pkill -9 -f ste_train 2>/dev/null") time.sleep(2) # 确认没有 ste_train 进程 import subprocess result = subprocess.run(["pgrep", "-f", "ste_train"], capture_output=True, text=True) if result.stdout.strip(): log(f"WARNING: {len(result.stdout.strip().split())} ste_train still running!") # 反复删除所有 ckpt 文件 (防止旧 worker 在删除后又 save) # 用 rm -f 而非 glob, 确保 model_dialogue.ste 被删 import glob for attempt in range(3): # 用 shell rm -f 强制删除 (glob 可能漏文件) os.system(f"rm -f {WORK_DIR}/ckpt_mp_* {WORK_DIR}/model_mp_*.ste " f"{WORK_DIR}/model_dialogue.ste {WORK_DIR}/merged_*.ste " f"{WORK_DIR}/model_*.ste 2>/dev/null") time.sleep(1) # 用 glob 检查是否还有残留 remaining = [] for pattern in [f"{WORK_DIR}/ckpt_mp_*", f"{WORK_DIR}/model_mp_*.ste", f"{WORK_DIR}/model_dialogue.ste", f"{WORK_DIR}/merged_*.ste", f"{WORK_DIR}/model_*.ste"]: remaining.extend(glob.glob(pattern)) if not remaining: log(f" attempt {attempt+1}: all ckpts deleted (0 remaining)") break log(f" attempt {attempt+1}: {len(remaining)} ckpts still exist, retry") time.sleep(2) if remaining: log(f"WARNING: {len(remaining)} ckpts could not be deleted: {remaining[:3]}") log("Fresh start: no resume, training from random init") # 4. 编译 (-O3 + SIMD + 对齐, 见 Makefile) — 强制重编译 # 用 subprocess 拿 make 的真实 rc (os.system + pipe 拿到的是 tail 的 rc) log("Compiling with -O3 SIMD + aligned alloc...") import subprocess try: result = subprocess.run( ["make", "-C", WORK_DIR], capture_output=True, text=True, timeout=120 ) if result.returncode != 0: log(f"make stderr (last 500): {result.stderr[-500:]}") log(f"make stdout (last 500): {result.stdout[-500:]}") rc = result.returncode except Exception as e: log(f"make exception: {e}") rc = 1 if rc != 0 or not os.path.exists(f"{WORK_DIR}/ste_train"): log(f"FATAL: make failed (rc={rc}), cannot start workers") return log("Compile OK") # 5. 启动多进程训练 start_workers() log(f"Training started: {N_WORKERS} workers × OMP=1") # 6. 后台监控: 每 500 步触发 auto-merge threading.Thread(target=upload_monitor, daemon=True).start() # ============================================================ # 上传 + 双 Space merge 监控 # ============================================================ def upload_monitor(): """每 500 步: 停 workers → merge 本地 → 上传 → 下载对方 → merge → 重启.""" api = HfApi(token=HF_TOKEN) if HF_TOKEN else None last_merged_step = 0 SPACE_ID = os.environ.get("LAL_SPACE_ID", "1") OTHER_SPACE_CKPT = f"ckpts/space{2 if SPACE_ID == '1' else 1}_ckpt.ste" MY_CKPT_NAME = f"ckpts/space{SPACE_ID}_ckpt.ste" while True: time.sleep(60) try: if not os.path.exists(TRAIN_LOG): continue with open(TRAIN_LOG) as f: lines = f.readlines() # 找最新的 CKPT saved 行 (worker 0) for line in reversed(lines): if "CKPT] saved" in line and "step" in line: step = int(line.split("step")[1].split("(")[0].strip()) if step > last_merged_step and step % UPLOAD_INTERVAL == 0 and step > 0: last_merged_step = step log(f"=== Auto-merge cycle triggered at step {step} ===") # 1. 停所有 worker stop_workers() # 2. merge 本 Space 多 worker ckpt → model_dialogue.ste merge_local_workers() # 3. 上传自己的 ckpt if api and os.path.exists(CKPT_PATH): log(f"Uploading {MY_CKPT_NAME} (step {step})...") api.upload_file( path_or_fileobj=CKPT_PATH, path_in_repo=MY_CKPT_NAME, repo_id=CKPT_REPO, repo_type="model" ) log("Upload done") # 4. 下载另一个 Space 的 ckpt (重试 6 次, 每次 60s = 共 6 分钟) # v21 fresh start 保护: 只在 step >= 500 时才 cross-merge # (避免 merge 旧污染 ckpt, 如 step < 500 时不下载对方) other_path = None if step < 500: log(f"Cross-Space merge skipped at step {step} < 500 (fresh start protection)") else: for attempt in range(6): try: other_path = hf_hub_download( repo_id=CKPT_REPO, filename=OTHER_SPACE_CKPT, repo_type="model" ) # 检查 ckpt 是否是新的 (修改时间在 30 分钟内) import os as _os mtime = _os.path.getmtime(other_path) age_s = time.time() - mtime if age_s < 1800: # 30 分钟内上传的 log(f"Downloaded other Space's ckpt: {OTHER_SPACE_CKPT} (age={age_s:.0f}s)") break else: log(f" attempt {attempt+1}: ckpt too old (age={age_s:.0f}s), retry in 60s") other_path = None except Exception as e: log(f" attempt {attempt+1}: not yet uploaded ({str(e)[:80]})") other_path = None if attempt < 5: time.sleep(60) else: log(f"Cross-Space merge skipped after 6 retries, continuing alone") # 5. 双 Space merge (只在下载成功时) if other_path: merged = f"{WORK_DIR}/merged_cross.ste" os.system( f"{WORK_DIR}/ste_train --merge {merged} " f"{CKPT_PATH} {other_path} 2>&1 | tail -3" ) if os.path.exists(merged) and os.path.getsize(merged) > 100_000_000: os.system(f"cp {merged} {CKPT_PATH}") os.remove(merged) log(f"Cross-Space merge done (step {step})") else: log("Cross-Space merge failed, keeping own ckpt") # 6. 重启多进程训练 start_workers() log(f"Workers resumed after merge cycle") break except Exception as e: log(f"monitor error: {e}") # ============================================================ # Gradio UI # ============================================================ def get_progress(): """获取 worker 0 的训练进度.""" header = f"**Multi-process: {N_WORKERS} workers**\n\n" # 先尝试 train.log (worker 0) for log_file in [TRAIN_LOG, f"{WORK_DIR}/worker_1.log"]: try: with open(log_file) as f: lines = f.readlines()[-20:] if lines and any("step" in l for l in lines): return header + "```\n" + "".join(lines) + "\n```" except: pass # 都没有, 显示 err 文件帮助 debug for err_file in [f"{WORK_DIR}/worker_0.err", f"{WORK_DIR}/train.err"]: try: with open(err_file) as f: err = f.read()[-500:] if err: return header + f"**worker_0.err**:\n```\n{err}\n```" except: pass return header + "Starting... (workers booting)" def get_worker_status(): """列出所有 worker 的最近一行日志.""" out = [] for i in range(N_WORKERS): log_file = f"{WORK_DIR}/train.log" if i == 0 else f"{WORK_DIR}/worker_{i}.log" err_file = f"{WORK_DIR}/worker_{i}.err" last_line = "no log" try: with open(log_file) as f: lines = [l for l in f.readlines() if l.strip()] if lines: last_line = lines[-1].strip()[:200] except: pass # 如果 log 空但 err 有内容, 显示 err if last_line == "no log": try: with open(err_file) as f: err = f.read().strip() if err: last_line = f"ERR: {err[-200:]}" except: pass out.append(f"**worker {i}**: `{last_line}`") return "\n".join(out) def run_inference(prompt): """运行推理 (先 merge workers → 推理).""" if not prompt: prompt = "你好" # 多进程下需要先 merge 才能 inference if not os.path.exists(CKPT_PATH) and os.path.exists(f"{WORK_DIR}/model_mp_0.ste"): merge_local_workers() prompt_file = f"{WORK_DIR}/prompt.txt" with open(prompt_file, "w") as f: f.write(prompt) try: result = subprocess.run( [f"{WORK_DIR}/ste_train", "--diagnose-only", "--prompt-file", prompt_file], capture_output=True, text=True, cwd=WORK_DIR, timeout=300 ) output = result.stdout + "\n" + result.stderr except subprocess.TimeoutExpired: output = "TIMEOUT (>300s)" return f"**Prompt**: {prompt}\n\n```\n{output}\n```" def manual_merge(): """手动触发 merge (UI 按钮).""" stop_workers() ok = merge_local_workers() start_workers() return f"Merge {'OK' if ok else 'FAILED'}, workers restarted" # 后台启动 threading.Thread(target=setup, daemon=True).start() @spaces.GPU def noop(): return "ok" with gr.Blocks() as demo: gr.Markdown(f"# LAL Training Worker (Multi-process: {N_WORKERS}×OMP=1)") with gr.Tab("训练进度"): progress_out = gr.Markdown(value="Starting...") refresh_btn = gr.Button("刷新进度") refresh_btn.click(get_progress, outputs=progress_out) demo.load(get_progress, outputs=progress_out) with gr.Tab("Worker 状态"): worker_out = gr.Markdown() wbtn = gr.Button("查看 Workers") wbtn.click(get_worker_status, outputs=worker_out) with gr.Tab("推理测试"): prompt_input = gr.Textbox(label="Prompt", value="什么是火") infer_btn = gr.Button("生成") infer_output = gr.Markdown() infer_btn.click(run_inference, inputs=prompt_input, outputs=infer_output) with gr.Tab("手动操作"): merge_btn = gr.Button("手动 Merge Workers") merge_out = gr.Markdown() merge_btn.click(manual_merge, outputs=merge_out) demo.launch(server_name="0.0.0.0", server_port=7860)