import os import torch from fastapi import FastAPI from pydantic import BaseModel, Field from comet import load_from_checkpoint from huggingface_hub import snapshot_download, HfApi # ========================================================== # 🚀 Configuração da API # ========================================================== app = FastAPI( title="XCOMET-XL API", version="1.5.0", description="API para avaliação de traduções usando Unbabel/XCOMET-XL, " "compatível com campos 'source', 'target' e 'human_translation_ref'." ) MODEL_NAME = "Unbabel/XCOMET-XL" HF_TOKEN = os.environ.get("HF_TOKEN") # defina nas Secrets do Space SPACE_REPO_ID = os.environ.get("SPACE_REPO_ID", "nairut/comet-xxl") # Diretório de cache local (dentro do espaço, permitido para escrita) MODEL_DIR = os.path.join(os.path.dirname(__file__), "model") MODEL_CKPT = os.path.join(MODEL_DIR, "checkpoints", "model.ckpt") # ========================================================== # ⚙️ Função auxiliar: baixa e persiste o modelo # ========================================================== def ensure_model_persisted_once(): """ Faz o download do modelo XCOMET-XL para ./model (caso ainda não exista) e tenta commitar essa pasta no próprio Space, para persistência. """ if os.path.exists(MODEL_CKPT): print(f"✅ Modelo já existe em {MODEL_CKPT}. Pulando download.") return print("🔽 Baixando snapshot do modelo para ./model ...") snapshot_download( repo_id=MODEL_NAME, token=HF_TOKEN, local_dir=MODEL_DIR, local_dir_use_symlinks=False ) assert os.path.exists(MODEL_CKPT), f"Checkpoint não encontrado: {MODEL_CKPT}" # tenta persistir no próprio Space (opcional) try: print("⬆️ Enviando pasta 'model/' para o repositório do Space ...") api = HfApi(token=HF_TOKEN) api.upload_folder( repo_id=SPACE_REPO_ID, repo_type="space", folder_path=MODEL_DIR, path_in_repo="model", commit_message="Persistência automática do modelo XCOMET-XL" ) print("✅ Modelo persistido no Space.") except Exception as e: print(f"⚠️ Falha ao persistir modelo no Space: {e}") print(" (prosseguindo com o modelo local para esta sessão)") # ========================================================== # 📦 Inicialização do modelo # ========================================================== ensure_model_persisted_once() print(f"📂 Carregando modelo de {MODEL_CKPT} ...") model = load_from_checkpoint(MODEL_CKPT) print("✅ Modelo XCOMET-XL carregado com sucesso!") USE_GPU = 1 if torch.cuda.is_available() else 0 print(f"⚙️ GPU detectada: {'sim' if USE_GPU else 'não'}") # ========================================================== # 🧠 Estrutura dos dados de entrada # ========================================================== class TranslationPair(BaseModel): source: str = Field(alias="source", description="Texto original") target: str = Field(alias="target", description="human-translation") machine_translation: str = Field(alias="machine_translation", description="human-translation") class Config: allow_population_by_field_name = True # ========================================================== # 🔧 Função utilitária # ========================================================== def prepare_data(pairs: list[TranslationPair]): """ Converte lista de TranslationPair no formato esperado pelo COMET: [{"src": ..., "mt": ..., "ref": ...}, ...] """ data = [] for p in pairs: src = str(p.source).strip() mt = str(p.machine_translation).strip() ref = str(p.target).strip() item = {"src": src, "mt": mt, "ref":ref} data.append(item) return data # ========================================================== # 🌐 Endpoints # ========================================================== @app.get("/") def root(): return { "message": "🚀 XCOMET-XL API ativa e pronta para uso!", "gpu_enabled": bool(USE_GPU), "available_endpoints": ["/score", "/score_batch"] } @app.post("/score") def score_single(pair: TranslationPair): """ Avalia um único par de tradução (source → target) com COMET-XL. """ try: data = [{ "src": str(pair.source), "mt": str(pair.target), "ref": str(pair.machine_rtanslation) }] output = model.predict(data, batch_size=8, gpus=USE_GPU) return { "system_score": getattr(output, "system_score", None), "segment_scores": getattr(output, "scores", None), "metadata": getattr(output, "metadata", None) } except Exception as e: print(f"❌ Erro em /score: {e}") return {"error": str(e)} @app.post("/score_batch") def score_batch(pairs: list[TranslationPair]): """ Avalia múltiplos pares de tradução em lote (batch). """ try: data = prepare_data(pairs) print(f"📊 Lote recebido: {len(data)} pares válidos") # reduz batch_size para evitar estouro de VRAM output = model.predict(data, batch_size=8, gpus=USE_GPU) return { "system_score": getattr(output, "system_score", None), "segment_scores": getattr(output, "scores", None), "metadata": getattr(output, "metadata", None) } except Exception as e: print(f"❌ Erro no batch: {e}") return {"error": str(e)} # ========================================================== # ▶️ Execução local (para debug) # ========================================================== if __name__ == "__main__": import uvicorn uvicorn.run(app, host="0.0.0.0", port=7860)