#!/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()