FloorMaterialAnalyzer/tools/model_server/server.py
2026-07-27 11:03:51 +08:00

137 lines
3.6 KiB
Python

import base64
import io
import os
from flask import Flask, jsonify, request
from PIL import Image
app = Flask(__name__)
app.config["MAX_CONTENT_LENGTH"] = 50 * 1024 * 1024
DEVICE = os.getenv("MODEL_DEVICE", "cuda")
DINO_MODEL = os.getenv("DINO_MODEL", "facebook/dinov2-large")
CLIP_MODEL = os.getenv("CLIP_MODEL", "ViT-L-14")
CLIP_PRETRAINED = os.getenv("CLIP_PRETRAINED", "laion2b_s32b_b82k")
_torch = None
_dino_processor = None
_dino_model = None
_clip_model = None
_clip_preprocess = None
def torch():
global _torch, DEVICE
if _torch is None:
import torch as torch_mod
_torch = torch_mod
if DEVICE == "cuda" and not _torch.cuda.is_available():
DEVICE = "cpu"
return _torch
def load_dino():
global _dino_processor, _dino_model
if _dino_model is None:
from transformers import AutoImageProcessor, AutoModel
_dino_processor = AutoImageProcessor.from_pretrained(DINO_MODEL)
_dino_model = AutoModel.from_pretrained(DINO_MODEL).to(DEVICE).eval()
return _dino_processor, _dino_model
def load_clip():
global _clip_model, _clip_preprocess
if _clip_model is None:
import open_clip
_clip_model, _, _clip_preprocess = open_clip.create_model_and_transforms(
CLIP_MODEL, pretrained=CLIP_PRETRAINED
)
_clip_model = _clip_model.to(DEVICE).eval()
return _clip_model, _clip_preprocess
def decode_image(payload):
raw = payload.get("image_base64")
if not raw:
raise ValueError("image_base64 is required")
data = base64.b64decode(raw)
return Image.open(io.BytesIO(data)).convert("RGB")
def normalize_tensor(vec):
t = torch()
vec = vec.float()
vec = vec / vec.norm(dim=-1, keepdim=True).clamp_min(1e-12)
return vec.detach().cpu().numpy()[0].astype("float32").tolist()
def dinov2_embedding(img):
t = torch()
processor, model = load_dino()
inputs = processor(images=img, return_tensors="pt")
inputs = {k: v.to(DEVICE) for k, v in inputs.items()}
with t.no_grad():
out = model(**inputs)
if getattr(out, "pooler_output", None) is not None:
vec = out.pooler_output
else:
vec = out.last_hidden_state[:, 0]
values = normalize_tensor(vec)
return {
"name": "dinov2_texture",
"model": "dinov2-large",
"dimension": len(values),
"normalize": True,
"provider": "model_server",
"generator": DINO_MODEL,
"values": values,
}
def clip_embedding(img):
t = torch()
model, preprocess = load_clip()
image_tensor = preprocess(img).unsqueeze(0).to(DEVICE)
with t.no_grad():
vec = model.encode_image(image_tensor)
values = normalize_tensor(vec)
return {
"name": "clip_visual",
"model": "ViT-L/14",
"dimension": len(values),
"normalize": True,
"provider": "model_server",
"generator": f"open_clip:{CLIP_MODEL}:{CLIP_PRETRAINED}",
"values": values,
}
@app.get("/health")
def health():
return jsonify(
{
"status": "ok",
"device": DEVICE,
"dino_model": DINO_MODEL,
"clip_model": CLIP_MODEL,
"clip_pretrained": CLIP_PRETRAINED,
}
)
@app.post("/embeddings")
def embeddings():
try:
img = decode_image(request.get_json(force=True))
return jsonify({"vectors": [dinov2_embedding(img), clip_embedding(img)]})
except Exception as exc:
return jsonify({"error": str(exc)}), 500
if __name__ == "__main__":
port = int(os.getenv("MODEL_SERVER_PORT", "5200"))
app.run(host="0.0.0.0", port=port)