duovlm-40m-v1 / code /webapp.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
4.07 kB
#!/usr/bin/env python3
"""DuoVLM-40M 交互页面(本地服务,浏览器用)
python webapp.py # 默认 0.0.0.0:8000
python webapp.py --port 8080
# 然后浏览器打开 http://localhost:8000
纯本地:模型在内存里只载一次;上传的图片走 base64 JSON,不经磁盘。
"""
from __future__ import annotations
import argparse
import base64
import io
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent))
from fastapi import FastAPI, HTTPException # noqa: E402
from fastapi.responses import HTMLResponse, JSONResponse # noqa: E402
from PIL import Image # noqa: E402
from pydantic import BaseModel # noqa: E402
HERE = Path(__file__).resolve().parent
PKG = HERE.parent
SAMPLES = PKG / "protocol" / "images"
app = FastAPI(title="DuoVLM-40M")
_vlm = None
def get_vlm():
global _vlm
if _vlm is None:
from duovlm_infer import DuoVLMInfer
_vlm = DuoVLMInfer()
return _vlm
class AskReq(BaseModel):
image_b64: str | None = None
sample: str | None = None
question: str = ""
blind: bool = False
max_new: int = 12
ngram: int = 3
rep_penalty: float = 1.0
temperature: float = 0.0
loop_break: bool = True
def _load_image(req: AskReq):
if req.sample:
p = SAMPLES / Path(req.sample).name
if not p.is_file():
raise HTTPException(404, f"包内没有样图 {req.sample}")
return p
if req.image_b64:
raw = base64.b64decode(req.image_b64.split(",", 1)[-1])
im = Image.open(io.BytesIO(raw)).convert("RGB")
if max(im.size) > 1400:
im.thumbnail((1400, 1400))
return im
return None
@app.get("/", response_class=HTMLResponse)
def index() -> str:
return (HERE / "index.html").read_text(encoding="utf-8")
@app.get("/api/health")
def health() -> JSONResponse:
v = get_vlm()
return JSONResponse({"ok": True, "params": v.n_params, "step": v.step, "device": v.dev,
"clip": str(v.clip_source), "extra": v.extra})
@app.get("/api/samples")
def samples() -> JSONResponse:
if not SAMPLES.is_dir():
return JSONResponse({"samples": []})
return JSONResponse({"samples": sorted(p.name for p in SAMPLES.glob("*.jpg"))[:12]})
@app.get("/api/sample/{name}")
def sample_image(name: str):
from fastapi.responses import FileResponse
p = SAMPLES / Path(name).name
if not p.is_file():
raise HTTPException(404, "包内没有这张样图")
return FileResponse(p, media_type="image/jpeg")
@app.post("/api/ask")
def ask(req: AskReq) -> JSONResponse:
v = get_vlm()
img = _load_image(req)
if img is None:
raise HTTPException(400, "需要 image_b64 或 sample")
kw = dict(max_new=req.max_new, ngram=req.ngram, rep_penalty=req.rep_penalty,
temperature=req.temperature, loop_break=req.loop_break)
out = v.ask(img, req.question, **kw)
if req.blind:
out["blind_answer"] = v.ask(img, req.question, blind=True, **kw)["answer"]
out["question"] = req.question or "Render a clear and concise summary of the photo."
return JSONResponse(out)
def main() -> None:
import socket
import uvicorn
ap = argparse.ArgumentParser()
ap.add_argument("--host", default="0.0.0.0")
ap.add_argument("--port", type=int, default=8000)
ap.add_argument("--no-warmup", action="store_true")
a = ap.parse_args()
if not a.no_warmup:
print("载入模型(首次约 5~15 秒)...", flush=True)
get_vlm()
print("模型就绪。", flush=True)
try:
ip = socket.gethostbyname(socket.gethostname())
except Exception:
ip = "127.0.0.1"
print(f"\n 浏览器打开: http://localhost:{a.port}\n"
f" 备用地址: http://{ip}:{a.port}\n", flush=True)
uvicorn.run(app, host=a.host, port=a.port, log_level="warning")
if __name__ == "__main__":
main()