import os import asyncio from fastapi import FastAPI from pydantic import BaseModel from huggingface_hub import snapshot_download from transformers import AutoTokenizer, AutoModelForCausalLM import torch import uuid app = FastAPI() MODEL_ID = "google/gemma-4-e4b-it" LOCAL_DIR = "/data/gemma4" HF_TOKEN = os.getenv("HF_TOKEN") # Kolejka FIFO queue = asyncio.Queue(maxsize=20) jobs = {} # job_id: status/text class GenReq(BaseModel): prompt: str max_new: int = 150 def ensure_model(): if not os.path.exists(f"{LOCAL_DIR}/config.json"): print("Ściągam Gemmę do /data...") snapshot_download(MODEL_ID, local_dir=LOCAL_DIR, token=HF_TOKEN, local_dir_use_symlinks=False) print("Model w /data gotowy") print("Start, sprawdzam model...") ensure_model() print("Ładuję tokenizer i model...") tokenizer = AutoTokenizer.from_pretrained(LOCAL_DIR, local_files_only=True) model = AutoModelForCausalLM.from_pretrained( LOCAL_DIR, device_map="cpu", torch_dtype=torch.bfloat16, local_files_only=True ) print("Model załadowany") async def worker(): while True: job_id, prompt, max_new = await queue.get() jobs[job_id] = {"status": "running"} inputs = tokenizer(prompt, return_tensors="pt") outputs = model.generate(**inputs, max_new_tokens=max_new, do_sample=True, temperature=0.7) text = tokenizer.decode(outputs[0], skip_special_tokens=True) jobs[job_id] = {"status": "done", "text": text} queue.task_done() @app.on_event("startup") async def startup(): asyncio.create_task(worker()) @app.post("/generate") async def generate(req: GenReq): if queue.full(): return {"error": "Kolejka pełna"} job_id = str(uuid.uuid4()) jobs[job_id] = {"status": "queued"} await queue.put((job_id, req.prompt, req.max_new)) return {"job_id": job_id, "position": queue.qsize()} @app.get("/result/{job_id}") async def result(job_id: str): return jobs.get(job_id, {"status": "not_found"})