BobCodeur/wolof-generation-api
0
1import torch2from fastapi import FastAPI3from fastapi.middleware.cors import CORSMiddleware4from pydantic import BaseModel5from transformers import AutoModelForCausalLM, AutoTokenizer6from peft import PeftModel7 8BASE = "Qwen/Qwen2.5-0.5B"9ADAPTER = "BobCodeur/qwen2.5-0.5b-wolof"10 11# Chargement une seule fois au démarrage : base + adaptateur LoRA wolof12tokenizer = AutoTokenizer.from_pretrained(ADAPTER)13base = AutoModelForCausalLM.from_pretrained(BASE)14model = PeftModel.from_pretrained(base, ADAPTER)15model.eval()16 17app = FastAPI(title="Wolof Text Generation API")18 19# CORS ouvert : l'API est appelable depuis n'importe quel site web20app.add_middleware(21 CORSMiddleware,22 allow_origins=["*"],23 allow_methods=["*"],24 allow_headers=["*"],25)26 27 28class GenIn(BaseModel):29 text: str30 max_length: int = 8031 temperature: float = 0.632 33 34def generer(text: str, max_length: int, temperature: float) -> str:35 inputs = tokenizer(text, return_tensors="pt")36 with torch.no_grad():37 sortie = model.generate(38 **inputs,39 max_new_tokens=int(max_length),40 do_sample=True,41 temperature=float(temperature),42 top_p=0.9,43 repetition_penalty=1.15,44 )45 return tokenizer.decode(sortie[0], skip_special_tokens=True)46 47 48@app.get("/")49def root():50 return {51 "message": "Wolof Text Generation API",52 "model": ADAPTER,53 "endpoints": {54 "GET /health": "vérifie que le service est réveillé",55 "POST /generate": "corps JSON {text, max_length?, temperature?}",56 "GET /generate?text=...": "test rapide au navigateur",57 },58 }59 60 61@app.get("/health")62def health():63 return {"status": "ok"}64 65 66@app.post("/generate")67def generate_post(inp: GenIn):68 generated = generer(inp.text, inp.max_length, inp.temperature)69 return {"prompt": inp.text, "generated": generated}70 71 72@app.get("/generate")73def generate_get(text: str, max_length: int = 80, temperature: float = 0.6):74 generated = generer(text, max_length, temperature)75 return {"prompt": text, "generated": generated}76 