#!/usr/bin/env python3 """一键验证:拿到包的人跑这一条命令,就知道包是不是完整、能不能用、数字对不对。 python verify_package.py # 完整验证(含 100 条迷你协议复算,需下载视觉塔约 600MB) python verify_package.py --quick # 只查校验和 / 参数量 / 推理连通(不下载视觉塔) 判据只用**客观项**:文件校验和、参数量、权重结构、推理可跑通、同输入结果稳定。 "模型答对了几条"只打印、不参与 PASS/FAIL —— 小模型命中率本来就低,拿它当判据会把好包判死。 """ from __future__ import annotations import argparse import hashlib import json import re import sys import time from pathlib import Path HERE = Path(__file__).resolve().parent PKG = HERE.parent sys.path.insert(0, str(HERE)) EXPECT_PARAMS = 39_335_936 EXPECT_LLM = 38_023_680 EXPECT_CONNECTOR = 1_312_256 EXPECT_STEP = 2474 # 公布值(logs/eval_vqa_s3.json,1,812 条 test):迷你集只是同分布抽样,用于对照、不判失败 PUBLISHED_TEST_ACC = 0.2953 PUBLISHED_TEST_N = 1812 ok_all = True def check(name: str, ok: bool, detail: str = "", hard: bool = True) -> bool: global ok_all mark = "PASS" if ok else ("FAIL" if hard else "WARN") if not ok and hard: ok_all = False print(f" [{mark}] {name}" + (f" {detail}" if detail else ""), flush=True) return ok def sha256(p: Path) -> str: h = hashlib.sha256() with open(p, "rb") as f: for b in iter(lambda: f.read(1 << 20), b""): h.update(b) return h.hexdigest() def step_sums() -> None: print("\n[1/5] 文件校验和 (SHA256SUMS)") sf = PKG / "SHA256SUMS" if not sf.is_file(): check("SHA256SUMS 存在", False) return bad = miss = 0 lines = 0 for line in sf.read_text().splitlines(): m = re.match(r"^([0-9a-f]{64})\s+(.+)$", line.strip()) if not m: continue lines += 1 want, rel = m.group(1), m.group(2) p = PKG / rel if not p.is_file(): print(f" 缺文件: {rel}") miss += 1 elif sha256(p) != want: print(f" 校验不符: {rel}") bad += 1 check(f"校验 {lines} 个文件", bad == 0 and miss == 0, f"不符 {bad} / 缺失 {miss}" + ("" if bad == 0 and miss == 0 else " → 包被改过或传输损坏")) def step_weights() -> None: print("\n[2/5] 权重结构与参数量") w = PKG / "weights" / "duovlm-s3-final.pth" if not w.is_file(): check("权重文件存在", False, str(w)) return import torch d = torch.load(w, map_location="cpu") check("顶层键 llm/connector/step/extra", all(k in d for k in ("llm", "connector", "step", "extra")), str(list(d.keys()))) check("step == 公布值", int(d["step"]) == EXPECT_STEP, f"{int(d['step'])} (期望 {EXPECT_STEP})") llm = d["llm"] con = d["connector"] tied = torch.equal(llm["lm_head.weight"], llm["transformer.wte.weight"]) n_raw = sum(v.numel() for v in llm.values() if hasattr(v, "numel")) # lm_head 与词表 embedding 共享同一份存储 → 可训练参数量要扣掉一次 n_llm = n_raw - llm["lm_head.weight"].numel() if tied else n_raw n_con = sum(v.numel() for v in con.values() if hasattr(v, "numel")) check("LLM 参数量", n_llm == EXPECT_LLM, f"{n_llm:,} (期望 {EXPECT_LLM:,})" + (";权重表条目含共享的 lm_head,已按绑定关系扣除" if tied else "")) check("connector 参数量", n_con == EXPECT_CONNECTOR, f"{n_con:,} (期望 {EXPECT_CONNECTOR:,})") check("合计可训练参数", n_llm + n_con == EXPECT_PARAMS, f"{n_llm + n_con:,} (期望 {EXPECT_PARAMS:,})") check("权重绑定 lm_head == wte", tied, "lm_head 与词表 embedding 共享(省 4.19M 参数)") def step_infer(quick: bool) -> None: print("\n[3/5] 推理连通性") from duovlm_infer import DuoVLMInfer t0 = time.time() v = DuoVLMInfer(verbose=False) check("模型载入", True, f"{v.n_params:,} 参数 / {v.dev} / {time.time()-t0:.1f}s") check("参数量与权重自洽", v.n_params == EXPECT_PARAMS, f"{v.n_params:,}") imgs = sorted((PKG / "protocol" / "images").glob("*.jpg")) if not imgs: check("包内样图存在", False) return r = v.ask(str(imgs[0]), "Render a clear and concise summary of the photo.") check("看图问答可跑通", bool(r["answer"]), f"『{r['answer']}』 {r['ms']}ms source={r['source']}") r2 = v.ask(str(imgs[0]), "Render a clear and concise summary of the photo.") check("同输入两次结果一致(解码稳定)", r["answer"] == r2["answer"], f"『{r2['answer']}』") def step_protocol() -> tuple[int, int]: print("\n[4/5] 迷你协议复算(100 条,全部图都在包内)") from duovlm_infer import DuoVLMInfer pf = PKG / "protocol" / "mini_vqa_100.jsonl" if not pf.is_file(): check("迷你协议文件存在", False, str(pf)) return 0, 0 rows = [json.loads(x) for x in pf.read_text().splitlines() if x.strip()] v = DuoVLMInfer(verbose=False) hit = 0 out = [] t0 = time.time() for r in rows: img = PKG / r["image"] got = v.ask(str(img), r["question"], max_new=12, ngram=3)["answer"] good = int(norm(got) == norm(r["answer"])) hit += good out.append({**r, "pred": got, "correct": bool(good)}) acc = hit / max(len(rows), 1) out_path = Path.cwd() / "mini_vqa_100_result.json" # 写在包外,避免污染包与校验和 out_path.write_text( json.dumps({"n": len(rows), "hit": hit, "acc": round(acc, 4), "seconds": round(time.time() - t0, 1), "rows": out}, ensure_ascii=False, indent=1)) print(f" 本机复算: {hit}/{len(rows)} = {acc*100:.1f}% ({time.time()-t0:.0f}s,逐条结果见 {out_path})") print(f" 公布值 : {PUBLISHED_TEST_N} 条 test 上 {PUBLISHED_TEST_ACC*100:.2f}% " f"(本迷你集为同分布抽样,仅作对照)") check("复算流程可跑通(不看命中率)", len(out) == len(rows), f"{len(out)}/{len(rows)} 条完成") return hit, len(rows) def norm(s: str) -> str: s = str(s).lower().strip() for a, b in (("'s", " s"), ("n't", " nt")): s = s.replace(a, b) s = re.sub(r"[;,!?]", "", s) s = re.sub(r"\.", "", s) s = re.sub(r"\b(a|an|the)\b", " ", s) return " ".join(s.split()) def dep_check() -> None: """先查依赖,给一句能照做的报错,而不是一串 ModuleNotFoundError 栈。""" missing = [] for m in ("torch", "transformers", "litgpt", "numpy", "PIL"): try: __import__(m) except ImportError: missing.append(m) if missing: print(f"缺依赖: {', '.join(missing)}\n" f"当前解释器: {sys.executable}\n\n" f"请先装依赖(建议独立虚拟环境):\n" f" python -m venv .venv && . .venv/bin/activate\n" f" pip install -r requirements.txt\n" f"然后用它跑: .venv/bin/python code/verify_package.py") sys.exit(2) from importlib.metadata import PackageNotFoundError, version import torch import transformers try: lt = version("litgpt") except PackageNotFoundError: lt = "未安装" print(f"解释器: {sys.executable}") print(f" torch {torch.__version__} / transformers {transformers.__version__} / " f"litgpt {lt} / cuda {torch.cuda.is_available()}" + ("" if lt == "0.5.13" else " ⚠ 期望 litgpt==0.5.13,其它版本可能接口不兼容")) def main() -> None: ap = argparse.ArgumentParser() ap.add_argument("--quick", action="store_true", help="跳过迷你协议复算(不下载视觉塔)") a = ap.parse_args() print("=" * 78) print("DuoVLM-40M 分发包自检") print(f"包目录: {PKG}") print("=" * 78) dep_check() step_sums() step_weights() if not a.quick: step_infer(a.quick) step_protocol() else: print("\n[3/5] 推理连通性 —— --quick 跳过(需要视觉塔)") print("\n[5/5] 汇总") print("=" * 78) print(("全部客观判据通过:包完整、参数对、能跑。" if ok_all else "有检查未通过,见上面 [FAIL]。")) print("注:模型答对几条只是参考,不参与判定;要复算头条数字请用完整 VQAv2 test 集。") sys.exit(0 if ok_all else 1) if __name__ == "__main__": main()