Files
enginex-bi100-compat/test_engine.py

185 lines
9.6 KiB
Python
Raw 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 -*-
"""
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")
ENTRYPOINT = os.path.join(HERE, "entrypoint.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")
p = subprocess.run(["bash", "-n", "entrypoint.sh"], capture_output=True, text=True, cwd=HERE)
check("entrypoint.sh bash 语法正确", p.returncode == 0, p.stderr[:200])
shutil.copy2(ENTRYPOINT, os.path.join(tmp, "ep.sh"))
p1 = subprocess.run(["bash", "-n", "ep.sh"], capture_output=True, text=True, cwd=tmp)
check("entrypoint.sh 复制后语法仍正确", p1.returncode == 0, p1.stderr[:200])
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())