GemmaAPI / app.py
VedaX-Labs's picture
Update app.py
94273af verified
Raw
History Blame Contribute Delete
3.92 kB
import os
import torch
import uvicorn
from fastapi import FastAPI
from pydantic import BaseModel
from transformers import AutoTokenizer, AutoModelForCausalLM
# ============================================================
# CONFIG
# ============================================================
MODEL_PATH = "./gemma3-270m-merged"
CPU_THREADS = int(os.environ.get("CPU_THREADS", "8"))
torch.set_num_threads(CPU_THREADS)
torch.set_num_interop_threads(2)
print("=" * 70)
print(" GEMMA 3 270M CPU API SERVER")
print("=" * 70)
print(f"Model: {MODEL_PATH}")
print("Device: cpu")
print(f"CPU threads: {CPU_THREADS}")
# ============================================================
# LOAD TOKENIZER
# ============================================================
print("\nLoading tokenizer...")
tokenizer = AutoTokenizer.from_pretrained(
MODEL_PATH
)
print("Tokenizer loaded.")
# ============================================================
# LOAD MODEL
# ============================================================
print("\nLoading model...")
model = AutoModelForCausalLM.from_pretrained(
MODEL_PATH,
dtype=torch.float32
)
model.eval()
print("Model loaded successfully.")
# ============================================================
# FASTAPI
# ============================================================
app = FastAPI(
title="Gemma 3 270M API",
version="1.0.0",
description="CPU inference API for Gemma 3 270M"
)
# ============================================================
# REQUEST FORMAT
# ============================================================
class ChatRequest(BaseModel):
message: str
max_new_tokens: int = 128
temperature: float = 0.7
top_p: float = 0.9
# ============================================================
# HEALTH
# ============================================================
@app.get("/health")
def health():
return {
"status": "ok",
"model": "Gemma 3 270M",
"device": "cpu"
}
# ============================================================
# CHAT
# ============================================================
@app.post("/v1/chat")
def chat(request: ChatRequest):
messages = [
{
"role": "user",
"content": request.message
}
]
inputs = tokenizer.apply_chat_template(
messages,
tokenize=True,
add_generation_prompt=True,
return_tensors="pt"
)
with torch.inference_mode():
outputs = model.generate(
input_ids=inputs,
max_new_tokens=request.max_new_tokens,
temperature=request.temperature,
top_p=request.top_p,
do_sample=True,
pad_token_id=tokenizer.eos_token_id
)
input_length = inputs.shape[-1]
generated_tokens = outputs[0][input_length:]
response = tokenizer.decode(
generated_tokens,
skip_special_tokens=True
)
return {
"response": response
}
# ============================================================
# ROOT
# ============================================================
@app.get("/")
def root():
return {
"name": "Gemma 3 270M API",
"status": "online",
"endpoints": {
"chat": "POST /v1/chat",
"health": "GET /health",
"docs": "GET /docs"
}
}
# ============================================================
# START
# ============================================================
if __name__ == "__main__":
print("\n" + "=" * 70)
print("SERVER READY")
print("=" * 70)
print("API:")
print("POST /v1/chat")
print("\nHealth:")
print("GET /health")
print("\nSwagger:")
print("GET /docs")
print("\nStarting server...")
uvicorn.run(
app,
host="0.0.0.0",
port=7860,
log_level="info"
)