NaiBot_API / main.py
ugonna's picture
Update main.py
69d444b verified
Raw History Blame Contribute Delete
6.36 kB
import torch
import os
import uvicorn
from fastapi import FastAPI, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
import time
import uuid
import logging
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger(__name__)
class GenerateRequest(BaseModel):
prompt: str
max_tokens: int = 512
temperature: float = 0.7
top_p: float = 0.9
class GenerateResponse(BaseModel):
id: str
text: str
created: int
class ModelManager:
def __init__(self):
self.model = None
self.tokenizer = None
self.device = None
self.loading_status = "not_started"
self.model_id = os.environ.get("MODEL_ID", "ugonna/llama3.18B-Fine-tunedByUgo3")
def setup_device(self):
if torch.cuda.is_available():
self.device = torch.device("cuda")
logger.info(f"✅ Using GPU: {torch.cuda.get_device_name(0)}")
else:
self.device = torch.device("cpu")
logger.info("⚠️ Using CPU - this will be VERY slow for large models")
def load_model(self):
try:
self.loading_status = "loading"
self.setup_device()
logger.info(f"📥 Loading tokenizer from {self.model_id}...")
self.tokenizer = AutoTokenizer.from_pretrained(self.model_id)
if self.tokenizer.pad_token is None:
self.tokenizer.pad_token = self.tokenizer.eos_token
logger.info(f"📥 Loading model from {self.model_id} with quantization...")
# Use 8-bit quantization to reduce memory (requires bitsandbytes)
bnb_config = BitsAndBytesConfig(
load_in_8bit=True,
llm_int8_threshold=6.0
)
self.model = AutoModelForCausalLM.from_pretrained(
self.model_id,
quantization_config=bnb_config if torch.cuda.is_available() else None,
device_map="auto" if torch.cuda.is_available() else None,
torch_dtype=torch.float16 if torch.cuda.is_available() else torch.float32,
low_cpu_mem_usage=True,
trust_remote_code=True
)
if not torch.cuda.is_available():
self.model = self.model.to(self.device)
self.model.eval()
self.loading_status = "loaded"
logger.info("✅ Model loaded successfully!")
return True
except Exception as e:
self.loading_status = "failed"
logger.error(f"❌ Failed to load model: {e}")
logger.error("This model may be too large for the free tier. Consider using a smaller model.")
return False
def generate(self, prompt: str, max_tokens: int = 512, temperature: float = 0.7, top_p: float = 0.9):
if self.model is None:
raise Exception("Model not loaded")
start_time = time.time()
inputs = self.tokenizer(prompt, return_tensors="pt", truncation=True, max_length=512)
if torch.cuda.is_available():
inputs = {k: v.cuda() for k, v in inputs.items()}
else:
inputs = {k: v.to(self.device) for k, v in inputs.items()}
with torch.no_grad():
outputs = self.model.generate(
**inputs,
max_new_tokens=min(max_tokens, 200), # Limit for CPU
temperature=temperature,
top_p=top_p,
do_sample=True,
pad_token_id=self.tokenizer.pad_token_id,
eos_token_id=self.tokenizer.eos_token_id
)
generated_ids = outputs[0][inputs['input_ids'].shape[1]:]
generated_text = self.tokenizer.decode(generated_ids, skip_special_tokens=True)
elapsed = time.time() - start_time
logger.info(f"✅ Generated {len(generated_ids)} tokens in {elapsed:.2f}s")
return generated_text
model_manager = ModelManager()
app = FastAPI(title="NAI Bot API", version="1.0.0")
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
@app.on_event("startup")
async def startup_event():
import threading
thread = threading.Thread(target=model_manager.load_model, daemon=True)
thread.start()
@app.get("/")
async def root():
return {
"message": "NAI Bot API is running",
"docs": "/docs",
"status": model_manager.loading_status,
"model": model_manager.model_id,
"warning": "Large model on free tier may fail due to memory limits"
}
@app.get("/health")
async def health_check():
return {
"status": "healthy" if model_manager.model is not None else "failed",
"model_loaded": model_manager.model is not None,
"loading_status": model_manager.loading_status,
"device": str(model_manager.device) if model_manager.device else "unknown",
"model_id": model_manager.model_id
}
@app.post("/generate", response_model=GenerateResponse)
async def generate(request: GenerateRequest):
if model_manager.model is None:
raise HTTPException(
status_code=503,
detail=f"Model failed to load. Your 8B model is too large for free tier. Please switch to a smaller model like 'microsoft/DialoGPT-small'"
)
if not request.prompt:
raise HTTPException(status_code=400, detail="Prompt cannot be empty")
try:
generated_text = model_manager.generate(
prompt=request.prompt,
max_tokens=request.max_tokens,
temperature=request.temperature,
top_p=request.top_p
)
return GenerateResponse(
id=str(uuid.uuid4()),
text=generated_text,
created=int(time.time())
)
except Exception as e:
logger.error(f"Generation error: {e}")
raise HTTPException(status_code=500, detail=str(e))
if __name__ == "__main__":
port = int(os.environ.get("PORT", 7860))
uvicorn.run(app, host="0.0.0.0", port=port)