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)