unimernet-test / backend_unimer.py
abcd-hide's picture
Upload 2 files
361cea9 verified
Raw
History Blame Contribute Delete
2.62 kB
import base64
import io
import os
import torch
import uvicorn
import traceback
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel
from PIL import Image
from transformers import AutoModel, AutoTokenizer
from fastapi.middleware.cors import CORSMiddleware
# --- アプリケーション設定 ---
app = FastAPI()
# CORSを許可
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_methods=["*"],
allow_headers=["*"],
)
# --- モデルのロード ---
model = None
tokenizer = None
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# ローカルにクローンされたパスを指定します
model_path = "/app/unimernet_model"
if not os.path.exists(model_path):
# ローカル開発環境用(Dockerfile外)
model_path = "wanderkid/unimernet_small"
print(f"UniMERNet モデルをロード中... (source: {model_path}, device: {device})")
try:
# ローカルパスからのロード。trust_remote_code=True はカスタムモデルコードの実行に必要です。
tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
model = AutoModel.from_pretrained(model_path, trust_remote_code=True).to(device)
model.eval()
print("モデルのロードが正常に完了しました。")
except Exception as e:
print(f"【致命的エラー】モデルのロードに失敗しました: {e}")
traceback.print_exc()
# --- リクエスト型定義 ---
class PredictionRequest(BaseModel):
image: str
# --- エンドポイント ---
@app.post("/predict")
async def predict(request: PredictionRequest):
if model is None:
raise HTTPException(status_code=503, detail="Model is not ready. Please check server logs.")
try:
if "," in request.image:
header, encoded = request.image.split(",", 1)
else:
encoded = request.image
image_data = base64.b64decode(encoded)
image = Image.open(io.BytesIO(image_data)).convert("RGB")
with torch.no_grad():
latex_code = model.predict(image, tokenizer)
return {"latex": latex_code}
except Exception as e:
print(f"推論中にエラーが発生しました: {e}")
traceback.print_exc()
raise HTTPException(status_code=500, detail=str(e))
@app.get("/health")
async def health():
if model is not None:
return {"status": "ready", "device": str(device)}
else:
return {"status": "loading_error"}
if __name__ == "__main__":
uvicorn.run(app, host="0.0.0.0", port=8000)