diff --git a/test_engine.py b/test_engine.py new file mode 100644 index 0000000..32ea17d --- /dev/null +++ b/test_engine.py @@ -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": ["", "", ""], + "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") == {"": "", "": "", "": ""}, + 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": {"": 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 ", '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())