Files
enginex-mlu370-compat/preflight.py

274 lines
11 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
preflight.py —— ModelHub XC 引擎启动前兼容修补器(Iluvatar bi-100 · vLLM)
在 vLLM 启动**之前**读取模型目录,把「引擎镜像层」的结构性不兼容修掉。
判据全部来自真实失败日志(见 README.md「证据」一节)。
已合并的两处**社区已上线修复**(来自 EngineX-Sunrise/enginex-S2-vllm-fix-tokenizer,
dev=sunruoxi,2026-05-28 上线,已在曦望 S2 生产使用):
· extra_special_tokens: list -> dict
· tokenizer_class 归一化(TokenizersBackend / TiktokenTokenizer -> 正常类)
加上本引擎针对 Iluvatar bi-100 失败日志新增的三处:
· R2 模型无 chat_template -> 注入默认 Jinja 模板
· R4 architectures 是镜像未注册类名 -> --hf-overrides 别名映射
· R7 镜像缺 ixformer.contrib.vllm.layers -> 纯 PyTorch shim(经 PYTHONPATH)
输出:JSON {"extra_args": [...], "patches": [...], "overlay": "..."}
extra_args 由 entrypoint.sh 追加到 `vllm serve` 之后。
设计红线:
* 绝不写模型目录(平台挂载的 /model 可能只读);只用临时覆盖目录 + 命令行参数
* 失败安全:任何一步异常只记 warning 并继续,最坏按原命令跑
* 干净模型零改动:不加任何参数
"""
import argparse
import json
import os
import shutil
import sys
import tempfile
# ---------------------------------------------------------------------------
# R4:已证实的「引擎镜像未注册 -> 已注册」架构别名表。只收有实证的,不猜。
# ---------------------------------------------------------------------------
ARCH_ALIASES = {
# 信创模盒官方基线 README 明确要求改成这个(Qwen3.6 系列)
"Qwen3_5MoeForConditionalGeneration": "Qwen3_5MoeForCausalLM",
"Qwen3_5ForConditionalGeneration": "Qwen3_5ForCausalLM",
# Gemma3 早期类名(寒武纪 mlu370 引擎实测 KeyError: 'gemma3_text')
"Gemma3ForConditionalGeneration": "Gemma3ForCausalLM",
"Gemma3TextForCausalLM": "Gemma3ForCausalLM",
}
# 社区已上线的「坏 tokenizer_class」清单(来自 fix_tokenizer.py)
BAD_TOKENIZER_CLASSES = ("TokenizersBackend", "TiktokenTokenizer")
# tokenizer 相关文件:覆盖目录只需要这些
TOKENIZER_FILES = (
"tokenizer.json", "tokenizer_config.json", "special_tokens_map.json",
"vocab.json", "merges.txt", "tokenizer.model", "added_tokens.json",
"chat_template.jinja", "preprocessor_config.json",
)
# 权重文件后缀:永远不拷贝/软链进覆盖目录
WEIGHT_SUFFIXES = (".safetensors", ".bin", ".gguf", ".pt", ".pth", ".msgpack",
".h5", ".onnx", ".ckpt", ".npz", ".arrow")
# R2:最小可用 Jinja 聊天模板(仅当模型自身没有时注入)
DEFAULT_CHAT_TEMPLATE = (
"{%- for message in messages -%}\n"
"{{- '' + message['role'] + '\\n' + message['content'] + '\\n' -}}\n"
"{%- endfor -%}\n"
"{%- if add_generation_prompt -%}\n"
"{{- 'assistant\\n' -}}\n"
"{%- endif -%}\n"
)
def _log(msg):
sys.stderr.write("[preflight] " + msg + "\n")
sys.stderr.flush()
def _load_json(path):
try:
with open(path, "r", encoding="utf-8") as f:
return json.load(f)
except Exception as e: # noqa: BLE001
_log("WARN 读不到 %s: %s" % (path, e))
return None
def _dump_json(obj, path):
"""写 JSON 到覆盖目录;失败返回 False。
自己负责把父目录建出来(不依赖调用方先前建过)——10-03 端到端测试出现过
"覆盖目录建了但 tokenizer 子目录写不进去"导致 R3 静默失效,这里做成自洽。
"""
try:
parent = os.path.dirname(os.path.abspath(path))
if parent:
os.makedirs(parent, exist_ok=True)
with open(path, "w", encoding="utf-8") as f:
json.dump(obj, f, ensure_ascii=False, indent=1)
return True
except Exception as e: # noqa: BLE001
_log("WARN 写不了 %s: %s" % (path, e))
return False
def _build_overlay(model_dir, overlay):
"""把 tokenizer 相关文件镜像进 overlay:优先软链,软链不可用则硬拷贝。
权重文件一律不碰。"""
tok_dir = os.path.join(overlay, "tokenizer")
os.makedirs(tok_dir, exist_ok=True)
n_link = n_copy = 0
try:
names = os.listdir(model_dir)
except Exception as e: # noqa: BLE001
_log("WARN listdir 失败 %s: %s" % (model_dir, e))
return tok_dir, 0, 0
# 1) 白名单 tokenizer 文件
for name in TOKENIZER_FILES:
s = os.path.join(model_dir, name)
d = os.path.join(tok_dir, name)
if not os.path.isfile(s) or os.path.exists(d) or os.path.islink(d):
continue
try:
os.symlink(s, d)
n_link += 1
except Exception: # noqa: BLE001
try:
shutil.copy2(s, d)
n_copy += 1
except Exception as e: # noqa: BLE001
_log("WARN 拷贝失败 %s: %s" % (s, e))
# 2) 其余非权重文件也软链过去(chat_template.* / *.jinja 等),保证 vLLM 找得到
for name in names:
if name in TOKENIZER_FILES or name.lower().endswith(WEIGHT_SUFFIXES):
continue
s = os.path.join(model_dir, name)
d = os.path.join(tok_dir, name)
if os.path.exists(d) or os.path.islink(d):
continue
try:
os.symlink(s, d)
n_link += 1
except Exception: # noqa: BLE001
pass
return tok_dir, n_link, n_copy
def detect_tokenizer(model_dir):
"""复刻社区 detect_tokenizer.py 的判定:fast / sentencepiece / bpe / unknown"""
files = set()
try:
files = set(os.listdir(model_dir))
except Exception: # noqa: BLE001
pass
if "tokenizer.json" in files:
return "fast"
if "tokenizer.model" in files:
return "sentencepiece"
if "vocab.json" in files and "merges.txt" in files:
return "bpe"
return "unknown"
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--model", default=os.environ.get("MODEL_DIR", "/model"))
ap.add_argument("--api", default="", choices=["", "chat", "completion"])
ap.add_argument("--out", default="")
args = ap.parse_args()
model_dir = args.model
patches = []
extra = []
if not os.path.isdir(model_dir):
res = {"extra_args": [], "patches": ["skip: 模型目录不存在 " + model_dir],
"overlay": ""}
_emit(res, args.out)
return 0
tok_cfg = _load_json(os.path.join(model_dir, "tokenizer_config.json"))
cfg = _load_json(os.path.join(model_dir, "config.json"))
tok_type = detect_tokenizer(model_dir)
need_tok_fix = False
if isinstance(tok_cfg, dict):
# 社区已上线的两条判据
if isinstance(tok_cfg.get("extra_special_tokens"), list):
need_tok_fix = True
if tok_cfg.get("tokenizer_class") in BAD_TOKENIZER_CLASSES:
need_tok_fix = True
overlay = ""
tok_dir = ""
if need_tok_fix:
overlay = tempfile.mkdtemp(prefix="mhxc_compat_")
tok_dir, nlk, ncp = _build_overlay(model_dir, overlay)
cfg_path = os.path.join(tok_dir, "tokenizer_config.json")
fixed = dict(tok_cfg)
# (a) extra_special_tokens list -> dict(与社区上线版同形:token -> token)
if isinstance(fixed.get("extra_special_tokens"), list):
orig = fixed["extra_special_tokens"]
fixed["extra_special_tokens"] = {t: t for t in orig}
patches.append("R3: extra_special_tokens list(%d) -> dict" % len(orig))
# (b) tokenizer_class 归一化
oc = fixed.get("tokenizer_class")
if oc in BAD_TOKENIZER_CLASSES:
if tok_type == "sentencepiece":
fixed["tokenizer_class"] = "LlamaTokenizer"
else:
fixed["tokenizer_class"] = "PreTrainedTokenizerFast"
fixed["from_slow"] = False
fixed.pop("backend", None)
patches.append("R3b: tokenizer_class %s -> %s" % (oc, fixed["tokenizer_class"]))
if _dump_json(fixed, cfg_path):
extra += ["--tokenizer", tok_dir]
patches.append("覆盖 tokenizer 目录 %s(软链 %d / 拷贝 %d,权重未碰)" % (tok_dir, nlk, ncp))
else:
patches.append("R3: 想修但写覆盖目录失败,按原样跑")
# ---------------- R2:模型无 chat_template -> 注入默认模板 ---------------
if args.api != "completion":
has_tpl = isinstance(tok_cfg, dict) and bool(tok_cfg.get("chat_template"))
has_tpl_file = False
try:
has_tpl_file = any(f.startswith("chat_template") for f in os.listdir(model_dir))
except Exception: # noqa: BLE001
pass
if not has_tpl and not has_tpl_file:
if not overlay:
overlay = tempfile.mkdtemp(prefix="mhxc_compat_")
tpl = os.path.join(overlay, "default_chat_template.jinja")
try:
with open(tpl, "w", encoding="utf-8") as f:
f.write(DEFAULT_CHAT_TEMPLATE)
extra += ["--chat-template", tpl]
patches.append("R2: 模型无 chat_template,注入默认模板 %s" % tpl)
except Exception as e: # noqa: BLE001
patches.append("R2: 想注入模板但失败 %s" % e)
# ---------------- R4:architectures 未注册 -> --hf-overrides -------------
if isinstance(cfg, dict):
archs = cfg.get("architectures") or []
if isinstance(archs, str):
archs = [archs]
if archs and archs[0] in ARCH_ALIASES:
new = ARCH_ALIASES[archs[0]]
extra += ["--hf-overrides",
json.dumps({"architectures": [new]}, separators=(",", ":"))]
patches.append("R4: architectures %s -> %s" % (archs[0], new))
# ---------------- R7:缺模块 -> PYTHONPATH 注入 shims --------------------
shim_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "shims")
if os.path.isdir(shim_dir):
prev = os.environ.get("PYTHONPATH", "")
os.environ["PYTHONPATH"] = shim_dir + (os.pathsep + prev if prev else "")
patches.append("R7: PYTHONPATH += %s" % shim_dir)
_emit({"model": model_dir, "api": args.api, "extra_args": extra,
"patches": patches, "overlay": overlay, "tokenizer_override": tok_dir},
args.out)
_log("patches=%d %s" % (len(patches), " | ".join(patches) if patches else "(无需修补)"))
return 0
def _emit(res, out_path):
txt = json.dumps(res, ensure_ascii=False, indent=1)
if out_path:
try:
with open(out_path, "w", encoding="utf-8") as f:
f.write(txt)
except Exception as e: # noqa: BLE001
_log("WARN 写 %s 失败: %s" % (out_path, e))
sys.stdout.write(txt + "\n")
if __name__ == "__main__":
sys.exit(main())