File size: 4,068 Bytes
4afe981 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 | #!/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()
|