185 lines
9.6 KiB
Python
185 lines
9.6 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")
|
|||
|
|
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())
|