Spaces:
Runtime error
Runtime error
| 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() | |
| async def startup(): | |
| asyncio.create_task(worker()) | |
| 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()} | |
| async def result(job_id: str): | |
| return jobs.get(job_id, {"status": "not_found"}) |