lal-train-2 / app.py
gasschina's picture
fix: use HF for source download (github clone failing in ZeroGPU)
f21482f verified
Raw
History Blame Contribute Delete
18.7 kB
#!/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()
@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)