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()