duovlm-40m-v1 / code /verify_package.py
Duoia's picture
DuoVLM-40M v1: from-scratch 40M vision-language model (frozen CLIP + MiniPile-pretrained LM)
4afe981 verified
Raw History Blame Contribute Delete
8.69 kB
#!/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()