from fastapi import FastAPI, File, UploadFile, HTTPException from fastapi.middleware.cors import CORSMiddleware from pydantic import BaseModel from PIL import Image import io import numpy as np from .sentence import TextEmbedder from .dinov2 import DinoV2 from .llava_next import LLaVANextCaptioner app = FastAPI() app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_credentials=True, allow_methods=["*"], allow_headers=["*"], ) print("-Loading LLaVA-Next-") llava = LLaVANextCaptioner("cjell/llava-next") print("-Loading DINOv2-") dino = DinoV2("cjell/dinov2") print("-Loading text embedding model-") text_embedder = TextEmbedder("cjell/text") print("\n-All models loaded successfully-") def read_image(file: UploadFile) -> Image.Image: try: contents = file.file.read() img = Image.open(io.BytesIO(contents)).convert("RGB") return img except Exception as e: raise HTTPException(status_code=400, detail=f"Invalid image: {e}") class TextPayload(BaseModel): text: str @app.post("/llava") async def caption_image(file: UploadFile = File(...)): img = read_image(file) caption = llava.caption(img) return {"caption": caption} @app.post("/dino") async def dino_embedding(file: UploadFile = File(...)): img = read_image(file) embedding = dino.embed_image(img) return {"embedding": embedding.tolist()} @app.post("/embed") async def text_embedding(payload: TextPayload): text = payload.text.strip() if not text: raise HTTPException(status_code=400, detail="Text cannot be empty.") vec = text_embedder.embed(text) return {"embedding": vec.tolist()} @app.get("/") async def root(): return { "status": "ok", "models": { "llava": "cjell/llava-next", "dino": "cjell/dinov2", "embedder": "cjell/text" } }