enginex-mrv100-compat: test_engine.py

This commit is contained in:
2026-10-03 20:05:05 +08:00
parent e7cc7f4a4d
commit 6fcb6e2111

186
test_engine.py Normal file
View File

@@ -0,0 +1,186 @@
#!/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())