Spaces:
Running on Zero
Running on Zero
| #!/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/lalmodel-code" # 数据文件 repo | |
| CKPT_REPO = "gasschina/lal-ckpts" # ckpt 共享 repo (gasschina 有写权限) | |
| 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/lalmodel-code 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/lalmodel-code (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/lalmodel-code/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. 下载数据 (v21: 纯 wiki 数据, 无数学污染) | |
| if not os.path.exists(f"{WORK_DIR}/data/wiki_clean.bin"): | |
| try: | |
| p = hf_hub_download(repo_id=DATA_REPO, filename="data/wiki_clean.bin", repo_type="model") | |
| os.system(f"mkdir -p {WORK_DIR}/data && cp {p} {WORK_DIR}/data/") | |
| log("wiki_clean.bin downloaded (pure wiki data, no math pollution)") | |
| 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"space{2 if SPACE_ID == '1' else 1}_ckpt.ste" | |
| MY_CKPT_NAME = f"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() | |
| 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) | |