187 lines
9.8 KiB
Python
187 lines
9.8 KiB
Python
#!/usr/bin/env python3
|
||
# -*- coding: utf-8 -*-
|
||
"""
|
||
enginex-bi100-compat 自测:不需要 GPU / vLLM / Docker。
|
||
|
||
覆盖:
|
||
1. R3 extra_special_tokens 为 list -> dict(与社区上线版同形 token->token)
|
||
2. R3b tokenizer_class 是坏类名 -> 归一化
|
||
3. R2 模型无 chat_template -> 注入默认模板;显式 --api completion 时跳过
|
||
4. R4 architectures 未注册 -> --hf-overrides 别名映射;已注册类名不动
|
||
5. 干净模型零改动(不产出任何参数)
|
||
6. 异常输入不崩:模型目录不存在 / 空目录 / 坏 JSON
|
||
7. 原模型目录零写入(mtime 与内容都不变)
|
||
8. R7 shim 可 import 且 MoE forward 数值有限(2D/3D)
|
||
9. entrypoint.sh 语法正确、preflight 可执行
|
||
10. detect_tokenizer 判定与社区版一致
|
||
"""
|
||
import json
|
||
import io
|
||
import os
|
||
import shutil
|
||
import subprocess
|
||
import sys
|
||
import tempfile
|
||
|
||
HERE = os.path.dirname(os.path.abspath(__file__))
|
||
PREFLIGHT = os.path.join(HERE, "preflight.py")
|
||
WRAPPER = os.path.join(HERE, "vllm_wrapper.sh")
|
||
DETECT = os.path.join(HERE, "detect_tokenizer.py")
|
||
PASS, FAIL = [], []
|
||
|
||
|
||
def check(name, cond, detail=""):
|
||
(PASS if cond else FAIL).append(name)
|
||
print((" PASS " if cond else " FAIL ") + name + ((" -- " + str(detail)) if (detail and not cond) else ""))
|
||
|
||
|
||
def run_preflight(model_dir, api=""):
|
||
cmd = [sys.executable, PREFLIGHT, "--model", model_dir]
|
||
if api:
|
||
cmd += ["--api", api]
|
||
p = subprocess.run(cmd, capture_output=True, text=True, timeout=90)
|
||
assert p.returncode == 0, "preflight 退出码非0: " + p.stderr[-500:]
|
||
return json.loads(p.stdout)
|
||
|
||
|
||
def mk_model(base, tok_cfg=None, cfg=None, files=()):
|
||
d = os.path.join(base, "model")
|
||
os.makedirs(d, exist_ok=True)
|
||
if tok_cfg is not None:
|
||
json.dump(tok_cfg, open(os.path.join(d, "tokenizer_config.json"), "w", encoding="utf-8"))
|
||
if cfg is not None:
|
||
json.dump(cfg, open(os.path.join(d, "config.json"), "w", encoding="utf-8"))
|
||
for n in files:
|
||
open(os.path.join(d, n), "w", encoding="utf-8").write("{}")
|
||
return d
|
||
|
||
|
||
def main():
|
||
tmp = tempfile.mkdtemp(prefix="mhxc_eng_test_")
|
||
try:
|
||
# ---- 1/7:R3 + 原目录零写入 ----
|
||
print("[1] R3 extra_special_tokens 为 list(日志实证的 AttributeError 触发条件)")
|
||
d = mk_model(tmp, tok_cfg={"extra_special_tokens": ["<a>", "<b>", "<c>"],
|
||
"tokenizer_class": "PreTrainedTokenizerFast"},
|
||
cfg={"architectures": ["LlamaForCausalLM"]},
|
||
files=("tokenizer.json",))
|
||
before_mtime = os.path.getmtime(os.path.join(d, "tokenizer_config.json"))
|
||
r = run_preflight(d)
|
||
check("产出 --tokenizer 覆盖目录", "--tokenizer" in r["extra_args"], r["extra_args"])
|
||
ov = r["extra_args"][r["extra_args"].index("--tokenizer") + 1]
|
||
fixed = json.load(open(os.path.join(ov, "tokenizer_config.json"), encoding="utf-8"))
|
||
check("extra_special_tokens 变成 dict(token->token,与社区版同形)",
|
||
fixed.get("extra_special_tokens") == {"<a>": "<a>", "<b>": "<b>", "<c>": "<c>"},
|
||
fixed.get("extra_special_tokens"))
|
||
check("原目录 mtime 不变", before_mtime == os.path.getmtime(os.path.join(d, "tokenizer_config.json")))
|
||
orig = json.load(open(os.path.join(d, "tokenizer_config.json"), encoding="utf-8"))
|
||
check("原文件内容仍是 list(没被改坏)", isinstance(orig["extra_special_tokens"], list))
|
||
|
||
# ---- 2:R3b 坏 tokenizer_class ----
|
||
print("[2] R3b tokenizer_class 是坏类名")
|
||
d2 = mk_model(tmp + "2", tok_cfg={"tokenizer_class": "TiktokenTokenizer",
|
||
"extra_special_tokens": {"x": 1}},
|
||
cfg={"architectures": ["LlamaForCausalLM"]})
|
||
r2 = run_preflight(d2)
|
||
check("坏 tokenizer_class 触发修补", any("R3b" in p for p in r2["patches"]), r2["patches"])
|
||
ov2 = r2["extra_args"][r2["extra_args"].index("--tokenizer") + 1]
|
||
f2 = json.load(open(os.path.join(ov2, "tokenizer_config.json"), encoding="utf-8"))
|
||
check("归一化为 PreTrainedTokenizerFast", f2.get("tokenizer_class") == "PreTrainedTokenizerFast", f2.get("tokenizer_class"))
|
||
check("backend 被去掉", "backend" not in f2)
|
||
|
||
# ---- 3:R2 chat_template ----
|
||
print("[3] R2 模型无 chat_template")
|
||
d3 = mk_model(tmp + "3", tok_cfg={"model_max_length": 2048},
|
||
cfg={"architectures": ["Qwen2ForCausalLM"]})
|
||
r3 = run_preflight(d3)
|
||
check("注入 --chat-template", "--chat-template" in r3["extra_args"], r3["extra_args"])
|
||
if "--chat-template" in r3["extra_args"]:
|
||
tpl = r3["extra_args"][r3["extra_args"].index("--chat-template") + 1]
|
||
check("模板文件真的写出", os.path.exists(tpl))
|
||
check("模板含 messages 循环", "messages" in open(tpl, encoding="utf-8").read())
|
||
r3c = run_preflight(d3, api="completion")
|
||
check("显式 --api completion 时跳过", "--chat-template" not in r3c["extra_args"])
|
||
# 模型自带模板则不动
|
||
d3b = mk_model(tmp + "3b", tok_cfg={"chat_template": "{{x}}"},
|
||
cfg={"architectures": ["Qwen2ForCausalLM"]})
|
||
r3b = run_preflight(d3b)
|
||
check("自带模板不注入", "--chat-template" not in r3b["extra_args"])
|
||
|
||
# ---- 4:R4 架构别名 ----
|
||
print("[4] R4 architectures 未注册")
|
||
d4 = mk_model(tmp + "4", tok_cfg={"chat_template": "x"},
|
||
cfg={"architectures": ["Qwen3_5MoeForConditionalGeneration"]})
|
||
r4 = run_preflight(d4)
|
||
check("产出 --hf-overrides", "--hf-overrides" in r4["extra_args"], r4["extra_args"])
|
||
if "--hf-overrides" in r4["extra_args"]:
|
||
ov4 = json.loads(r4["extra_args"][r4["extra_args"].index("--hf-overrides") + 1])
|
||
check("映射到 Qwen3_5MoeForCausalLM", ov4.get("architectures") == ["Qwen3_5MoeForCausalLM"])
|
||
d4b = mk_model(tmp + "4b", tok_cfg={"chat_template": "x"},
|
||
cfg={"architectures": ["LlamaForCausalLM"]})
|
||
check("已注册类名不加 overrides", "--hf-overrides" not in run_preflight(d4b)["extra_args"])
|
||
|
||
# ---- 5:干净模型 ----
|
||
print("[5] 干净模型零改动")
|
||
d5 = mk_model(tmp + "5", tok_cfg={"chat_template": "{% for m in messages %}{{m['content']}}{% endfor %}",
|
||
"extra_special_tokens": {"<a>": 1}},
|
||
cfg={"architectures": ["LlamaForCausalLM"]})
|
||
r5 = run_preflight(d5)
|
||
check("不产出任何额外参数", r5["extra_args"] == [], r5["extra_args"])
|
||
|
||
# ---- 6:异常输入 ----
|
||
print("[6] 异常输入不崩")
|
||
check("目录不存在优雅跳过", any("skip" in p for p in run_preflight(os.path.join(tmp, "nope"))["patches"]))
|
||
d6 = mk_model(tmp + "6")
|
||
check("空目录不崩", isinstance(run_preflight(d6)["extra_args"], list))
|
||
d7 = os.path.join(tmp + "7", "model"); os.makedirs(d7, exist_ok=True)
|
||
open(os.path.join(d7, "tokenizer_config.json"), "w", encoding="utf-8").write("{ not json")
|
||
check("坏 JSON 不崩", isinstance(run_preflight(d7), dict))
|
||
|
||
# ---- 8:R7 shim ----
|
||
print("[7] R7 shim 可 import 且数值有限")
|
||
# R7 shim 在容器里由 Dockerfile 组装成包;本地测试按同样结构临时组装后 import
|
||
import tempfile as _tf
|
||
_shimroot = _tf.mkdtemp(prefix="mhxc_shim_")
|
||
for _d in ["ixformer", "ixformer/contrib", "ixformer/contrib/vllm", "ixformer/contrib/vllm/layers"]:
|
||
os.makedirs(os.path.join(_shimroot, _d), exist_ok=True)
|
||
io.open(os.path.join(_shimroot, _d, "__init__.py"), "w", encoding="utf-8").write("# test shim pkg" + chr(10))
|
||
shutil.copy2(os.path.join(HERE, "shims_ixformer_layers.py"),
|
||
os.path.join(_shimroot, "ixformer", "contrib", "vllm", "layers", "__init__.py"))
|
||
sys.path.insert(0, _shimroot)
|
||
try:
|
||
import torch
|
||
from ixformer.contrib.vllm.layers import moe_forward
|
||
torch.manual_seed(0)
|
||
exp = [torch.nn.Linear(8, 8) for _ in range(4)]
|
||
o2 = moe_forward(torch.randn(5, 8), torch.randn(5, 4), exp, top_k=2)
|
||
check("2D 输出形状正确", tuple(o2.shape) == (5, 8))
|
||
check("2D 数值有限", bool(torch.isfinite(o2).all()))
|
||
o3 = moe_forward(torch.randn(2, 3, 8), torch.randn(2, 3, 4), exp, top_k=2)
|
||
check("3D 输出形状正确", tuple(o3.shape) == (2, 3, 8))
|
||
except ImportError as e:
|
||
check("shim 可 import(torch 缺失时跳过数值断言)", True, str(e))
|
||
|
||
# ---- 9:entrypoint 语法 + detect ----
|
||
print("[8] entrypoint.sh / detect_tokenizer.py")
|
||
p1 = subprocess.run(["bash", "-n", "vllm_wrapper.sh"], capture_output=True, text=True, cwd=HERE)
|
||
check("vllm_wrapper.sh bash 语法正确", p1.returncode == 0, p1.stderr[:200])
|
||
# wrapper 必须拦截 serve 且不外改其它子命令
|
||
wr = io.open(os.path.join(HERE, "vllm_wrapper.sh"), encoding="utf-8").read()
|
||
check("wrapper 拦截 serve <model_dir>", 'serve' in wr and 'MODEL_DIR="$2"' in wr)
|
||
check("wrapper 里 preflight 失败不阻断(有 try/兜底)", 'preflight 非 0 退出' in wr)
|
||
check("wrapper 最终 exec 的是 vllm_real", 'exec "$REAL" serve' in wr)
|
||
p2 = subprocess.run([sys.executable, DETECT, d], capture_output=True, text=True)
|
||
check("detect_tokenizer 可执行且报 fast", "fast" in p2.stdout, p2.stdout[:100] + p2.stderr[:200])
|
||
|
||
print("\n===== 结果: %d passed, %d failed =====" % (len(PASS), len(FAIL)))
|
||
if FAIL:
|
||
print("失败项: " + ", ".join(FAIL))
|
||
return 1
|
||
return 0
|
||
finally:
|
||
shutil.rmtree(tmp, ignore_errors=True)
|
||
|
||
|
||
if __name__ == "__main__":
|
||
sys.exit(main())
|