ugonna commited on
Commit
aa09e62
·
verified ·
1 Parent(s): f54be41

Update main.py

Browse files
Files changed (1) hide show
  1. main.py +83 -56
main.py CHANGED
@@ -3,84 +3,94 @@ import os
3
  import uvicorn
4
  from fastapi import FastAPI, HTTPException
5
  from fastapi.middleware.cors import CORSMiddleware
6
- from pydantic import BaseModel, Field
7
- from transformers import AutoModelForCausalLM, AutoTokenizer, BitsAndBytesConfig
8
  import time
9
  import uuid
10
- from typing import Dict, Any
11
  import logging
12
 
13
  logging.basicConfig(level=logging.INFO)
14
  logger = logging.getLogger(__name__)
15
 
 
 
 
 
 
 
 
 
 
 
 
 
16
  # Model Manager
17
  class ModelManager:
18
  def __init__(self):
19
  self.model = None
20
  self.tokenizer = None
21
  self.device = None
22
- self.model_path = "/opt/render/project/src/models"
23
  self.loading_status = "not_started"
24
-
 
 
25
  def setup_device(self):
26
  if torch.cuda.is_available():
27
  self.device = torch.device("cuda")
28
- logger.info(f"Using GPU: {torch.cuda.get_device_name(0)}")
29
  else:
30
  self.device = torch.device("cpu")
31
- logger.info("Using CPU")
32
 
33
  def load_model(self):
34
  try:
35
  self.loading_status = "loading"
36
  self.setup_device()
37
 
38
- logger.info("Loading tokenizer...")
39
- self.tokenizer = AutoTokenizer.from_pretrained(self.model_path)
40
  if self.tokenizer.pad_token is None:
41
  self.tokenizer.pad_token = self.tokenizer.eos_token
42
 
43
- logger.info("Loading model...")
44
- if torch.cuda.is_available():
45
- self.model = AutoModelForCausalLM.from_pretrained(
46
- self.model_path,
47
- torch_dtype=torch.float16,
48
- device_map="auto",
49
- trust_remote_code=True
50
- )
51
- else:
52
- self.model = AutoModelForCausalLM.from_pretrained(
53
- self.model_path,
54
- torch_dtype=torch.float32,
55
- trust_remote_code=True
56
- ).to(self.device)
57
 
58
  self.loading_status = "loaded"
59
  logger.info("✅ Model loaded successfully!")
60
  return True
 
61
  except Exception as e:
62
  self.loading_status = "failed"
63
- logger.error(f"Failed to load model: {e}")
64
  return False
65
 
66
- def generate(self, prompt: str, **kwargs):
67
  if self.model is None:
68
  raise Exception("Model not loaded")
69
 
70
  start_time = time.time()
71
- inputs = self.tokenizer(prompt, return_tensors="pt", truncation=True, max_length=2048)
72
 
73
- if torch.cuda.is_available():
74
- inputs = {k: v.cuda() for k, v in inputs.items()}
75
- else:
76
- inputs = {k: v.to(self.device) for k, v in inputs.items()}
77
 
78
  with torch.no_grad():
79
  outputs = self.model.generate(
80
  **inputs,
81
- max_new_tokens=kwargs.get('max_tokens', 512),
82
- temperature=kwargs.get('temperature', 0.7),
83
- top_p=kwargs.get('top_p', 0.9),
84
  do_sample=True,
85
  pad_token_id=self.tokenizer.pad_token_id,
86
  eos_token_id=self.tokenizer.eos_token_id
@@ -89,14 +99,18 @@ class ModelManager:
89
  generated_ids = outputs[0][inputs['input_ids'].shape[1]:]
90
  generated_text = self.tokenizer.decode(generated_ids, skip_special_tokens=True)
91
 
92
- return {"text": generated_text}
 
 
 
93
 
94
  # Initialize model manager
95
  model_manager = ModelManager()
96
 
97
- # Create FastAPI app
98
- app = FastAPI(title="Model API", version="1.0.0")
99
 
 
100
  app.add_middleware(
101
  CORSMiddleware,
102
  allow_origins=["*"],
@@ -107,42 +121,55 @@ app.add_middleware(
107
 
108
  @app.on_event("startup")
109
  async def startup_event():
110
- """Load model on startup"""
111
  import threading
112
  thread = threading.Thread(target=model_manager.load_model, daemon=True)
113
  thread.start()
114
 
 
 
 
 
 
 
 
 
 
115
  @app.get("/health")
116
  async def health_check():
117
  return {
118
  "status": "healthy" if model_manager.model is not None else "loading",
119
  "model_loaded": model_manager.model is not None,
120
- "device": str(model_manager.device) if model_manager.device else "unknown"
 
 
121
  }
122
 
123
- @app.post("/generate")
124
- async def generate(request: dict):
125
  if model_manager.model is None:
126
- raise HTTPException(status_code=503, detail="Model not ready")
 
 
 
127
 
128
  try:
129
- result = model_manager.generate(
130
- prompt=request.get("prompt", ""),
131
- max_tokens=request.get("max_tokens", 512),
132
- temperature=request.get("temperature", 0.7)
 
 
 
 
 
 
 
133
  )
134
- return {
135
- "id": str(uuid.uuid4()),
136
- "text": result["text"],
137
- "created": int(time.time())
138
- }
139
  except Exception as e:
 
140
  raise HTTPException(status_code=500, detail=str(e))
141
 
142
- @app.get("/")
143
- async def root():
144
- return {"message": "Model API is running", "docs": "/docs"}
145
-
146
  if __name__ == "__main__":
147
- port = int(os.environ.get("PORT", 8000))
148
- uvicorn.run(app, host="0.0.0.0", port=port)
 
3
  import uvicorn
4
  from fastapi import FastAPI, HTTPException
5
  from fastapi.middleware.cors import CORSMiddleware
6
+ from pydantic import BaseModel
7
+ from transformers import AutoModelForCausalLM, AutoTokenizer
8
  import time
9
  import uuid
 
10
  import logging
11
 
12
  logging.basicConfig(level=logging.INFO)
13
  logger = logging.getLogger(__name__)
14
 
15
+ # Request/Response Models
16
+ class GenerateRequest(BaseModel):
17
+ prompt: str
18
+ max_tokens: int = 512
19
+ temperature: float = 0.7
20
+ top_p: float = 0.9
21
+
22
+ class GenerateResponse(BaseModel):
23
+ id: str
24
+ text: str
25
+ created: int
26
+
27
  # Model Manager
28
  class ModelManager:
29
  def __init__(self):
30
  self.model = None
31
  self.tokenizer = None
32
  self.device = None
 
33
  self.loading_status = "not_started"
34
+ # Use a model from Hugging Face Hub
35
+ self.model_id = os.environ.get("MODEL_ID", "microsoft/DialoGPT-small")
36
+
37
  def setup_device(self):
38
  if torch.cuda.is_available():
39
  self.device = torch.device("cuda")
40
+ logger.info(f"✅ Using GPU: {torch.cuda.get_device_name(0)}")
41
  else:
42
  self.device = torch.device("cpu")
43
+ logger.info("⚠️ Using CPU (slower for LLMs)")
44
 
45
  def load_model(self):
46
  try:
47
  self.loading_status = "loading"
48
  self.setup_device()
49
 
50
+ logger.info(f"📥 Loading tokenizer from {self.model_id}...")
51
+ self.tokenizer = AutoTokenizer.from_pretrained(self.model_id)
52
  if self.tokenizer.pad_token is None:
53
  self.tokenizer.pad_token = self.tokenizer.eos_token
54
 
55
+ logger.info(f"📥 Loading model from {self.model_id}...")
56
+
57
+ # Load with optimizations for CPU
58
+ self.model = AutoModelForCausalLM.from_pretrained(
59
+ self.model_id,
60
+ torch_dtype=torch.float32, # Use float32 for CPU
61
+ low_cpu_mem_usage=True,
62
+ trust_remote_code=True
63
+ )
64
+
65
+ # Move to CPU explicitly
66
+ self.model = self.model.to(self.device)
67
+ self.model.eval()
 
68
 
69
  self.loading_status = "loaded"
70
  logger.info("✅ Model loaded successfully!")
71
  return True
72
+
73
  except Exception as e:
74
  self.loading_status = "failed"
75
+ logger.error(f"❌ Failed to load model: {e}")
76
  return False
77
 
78
+ def generate(self, prompt: str, max_tokens: int = 512, temperature: float = 0.7, top_p: float = 0.9):
79
  if self.model is None:
80
  raise Exception("Model not loaded")
81
 
82
  start_time = time.time()
83
+ inputs = self.tokenizer(prompt, return_tensors="pt", truncation=True, max_length=1024)
84
 
85
+ # Move inputs to the same device as model
86
+ inputs = {k: v.to(self.device) for k, v in inputs.items()}
 
 
87
 
88
  with torch.no_grad():
89
  outputs = self.model.generate(
90
  **inputs,
91
+ max_new_tokens=max_tokens,
92
+ temperature=temperature,
93
+ top_p=top_p,
94
  do_sample=True,
95
  pad_token_id=self.tokenizer.pad_token_id,
96
  eos_token_id=self.tokenizer.eos_token_id
 
99
  generated_ids = outputs[0][inputs['input_ids'].shape[1]:]
100
  generated_text = self.tokenizer.decode(generated_ids, skip_special_tokens=True)
101
 
102
+ elapsed = time.time() - start_time
103
+ logger.info(f"✅ Generated {len(generated_ids)} tokens in {elapsed:.2f}s")
104
+
105
+ return generated_text
106
 
107
  # Initialize model manager
108
  model_manager = ModelManager()
109
 
110
+ # FastAPI app
111
+ app = FastAPI(title="NAI Bot API", version="1.0.0")
112
 
113
+ # CORS
114
  app.add_middleware(
115
  CORSMiddleware,
116
  allow_origins=["*"],
 
121
 
122
  @app.on_event("startup")
123
  async def startup_event():
124
+ """Load model in background on startup"""
125
  import threading
126
  thread = threading.Thread(target=model_manager.load_model, daemon=True)
127
  thread.start()
128
 
129
+ @app.get("/")
130
+ async def root():
131
+ return {
132
+ "message": "NAI Bot API is running",
133
+ "docs": "/docs",
134
+ "status": model_manager.loading_status,
135
+ "model": model_manager.model_id
136
+ }
137
+
138
  @app.get("/health")
139
  async def health_check():
140
  return {
141
  "status": "healthy" if model_manager.model is not None else "loading",
142
  "model_loaded": model_manager.model is not None,
143
+ "loading_status": model_manager.loading_status,
144
+ "device": str(model_manager.device) if model_manager.device else "unknown",
145
+ "model_id": model_manager.model_id
146
  }
147
 
148
+ @app.post("/generate", response_model=GenerateResponse)
149
+ async def generate(request: GenerateRequest):
150
  if model_manager.model is None:
151
+ raise HTTPException(status_code=503, detail=f"Model is still loading (status: {model_manager.loading_status}). Try again in a few seconds.")
152
+
153
+ if not request.prompt:
154
+ raise HTTPException(status_code=400, detail="Prompt cannot be empty")
155
 
156
  try:
157
+ generated_text = model_manager.generate(
158
+ prompt=request.prompt,
159
+ max_tokens=request.max_tokens,
160
+ temperature=request.temperature,
161
+ top_p=request.top_p
162
+ )
163
+
164
+ return GenerateResponse(
165
+ id=str(uuid.uuid4()),
166
+ text=generated_text,
167
+ created=int(time.time())
168
  )
 
 
 
 
 
169
  except Exception as e:
170
+ logger.error(f"Generation error: {e}")
171
  raise HTTPException(status_code=500, detail=str(e))
172
 
 
 
 
 
173
  if __name__ == "__main__":
174
+ port = int(os.environ.get("PORT", 7860))
175
+ uvicorn.run(app, host="0.0.0.0", port=port)