Singh commited on
Commit
36e1d48
Β·
verified Β·
1 Parent(s): dc00de0

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +42 -36
app.py CHANGED
@@ -11,7 +11,6 @@ from typing import Optional, Iterator
11
  logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s")
12
  logger = logging.getLogger(__name__)
13
 
14
- # ── Model config β€” opt-350m runs fine on CPU, no token needed ──
15
  MODEL_ID = os.getenv("MODEL_ID", "facebook/opt-350m")
16
  HF_TOKEN = os.getenv("HF_TOKEN", None)
17
  DEVICE = "cpu"
@@ -23,25 +22,16 @@ tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, token=HF_TOKEN)
23
  model = AutoModelForCausalLM.from_pretrained(
24
  MODEL_ID,
25
  torch_dtype=DTYPE,
26
- device_map="auto",
27
  token=HF_TOKEN,
28
  )
29
  model.eval()
30
  logger.info("Model ready.")
31
 
32
- api = FastAPI(title="Gemma 2B API")
33
- api.add_middleware(
34
- CORSMiddleware,
35
- allow_origins=["*"],
36
- allow_credentials=False,
37
- allow_methods=["*"],
38
- allow_headers=["*"],
39
- expose_headers=["*"],
40
- )
41
 
42
  class GenerateRequest(BaseModel):
43
  prompt: str = Field(..., min_length=1, max_length=4096)
44
- max_new_tokens: int = Field(default=256, ge=1, le=1024)
45
  temperature: float = Field(default=0.7, ge=0.01, le=2.0)
46
  top_p: float = Field(default=0.9, ge=0.0, le=1.0)
47
  top_k: int = Field(default=50, ge=0, le=200)
@@ -52,23 +42,16 @@ class GenerateRequest(BaseModel):
52
  def stream_tokens(req: GenerateRequest) -> Iterator[str]:
53
  try:
54
  if req.system_prompt:
55
- prompt = (
56
- f"<start_of_turn>system\n{req.system_prompt}<end_of_turn>\n"
57
- f"<start_of_turn>user\n{req.prompt}<end_of_turn>\n"
58
- f"<start_of_turn>model\n"
59
- )
60
  else:
61
- prompt = (
62
- f"<start_of_turn>user\n{req.prompt}<end_of_turn>\n"
63
- f"<start_of_turn>model\n"
64
- )
65
 
66
  inputs = tokenizer(
67
- prompt, return_tensors="pt", truncation=True, max_length=2048
68
- ).to(model.device)
69
 
70
  streamer = TextIteratorStreamer(
71
- tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=60.0
72
  )
73
 
74
  gen_kwargs = dict(
@@ -97,7 +80,8 @@ def stream_tokens(req: GenerateRequest) -> Iterator[str]:
97
 
98
  t.join()
99
  latency = (time.perf_counter() - start) * 1000
100
- yield f"data: {json.dumps({'done': True, 'total_tokens': token_count, 'latency_ms': round(latency,1)})}\n\n"
 
101
 
102
  except Exception as e:
103
  tb = traceback.format_exc()
@@ -105,17 +89,29 @@ def stream_tokens(req: GenerateRequest) -> Iterator[str]:
105
  yield f"data: {json.dumps({'error': str(e), 'traceback': tb})}\n\n"
106
 
107
 
108
- @api.get("/health")
 
 
 
 
 
 
 
 
 
 
 
 
109
  async def health():
110
  return {
111
- "status" : "ok",
112
- "model" : MODEL_ID,
113
- "device" : DEVICE,
114
- "gpu_memory_used_gb" : round(torch.cuda.memory_allocated() / 1e9, 2) if DEVICE == "cuda" else 0,
115
- "gpu_memory_total_gb" : round(torch.cuda.get_device_properties(0).total_memory / 1e9, 2) if DEVICE == "cuda" else 0,
116
  }
117
 
118
- @api.post("/generate/stream")
 
119
  async def generate_stream(req: GenerateRequest):
120
  return StreamingResponse(
121
  stream_tokens(req),
@@ -127,11 +123,13 @@ async def generate_stream(req: GenerateRequest):
127
  },
128
  )
129
 
130
- @api.post("/generate")
 
131
  async def generate(req: GenerateRequest):
132
  full_text = ""
133
  total_tokens = 0
134
  latency_ms = 0.0
 
135
  for chunk in stream_tokens(req):
136
  if not chunk.startswith("data: "):
137
  continue
@@ -146,6 +144,7 @@ async def generate(req: GenerateRequest):
146
  total_tokens += 1
147
  if "done" in data:
148
  latency_ms = data["latency_ms"]
 
149
  return {
150
  "generated_text" : full_text,
151
  "completion_tokens" : total_tokens,
@@ -154,6 +153,7 @@ async def generate(req: GenerateRequest):
154
  }
155
 
156
 
 
157
  def gradio_generate(prompt, system_prompt, max_new_tokens, temperature):
158
  req = GenerateRequest(
159
  prompt = prompt,
@@ -173,9 +173,9 @@ def gradio_generate(prompt, system_prompt, max_new_tokens, temperature):
173
  pass
174
 
175
 
176
- with gr.Blocks(title="Gemma 2B API") as demo:
177
  gr.Markdown(
178
- "## Gemma 2B β€” Streaming API\n"
179
  "Use `/api/generate/stream` or `/api/generate` from your backend.\n\n"
180
  "**Health check:** `/api/health`"
181
  )
@@ -184,7 +184,7 @@ with gr.Blocks(title="Gemma 2B API") as demo:
184
  sys_box = gr.Textbox(label="System prompt (optional)", lines=2)
185
  prompt_box = gr.Textbox(label="Prompt", lines=4, placeholder="Ask something...")
186
  with gr.Row():
187
- max_tok = gr.Slider(32, 1024, value=256, step=32, label="Max tokens")
188
  temp = gr.Slider(0.01, 2.0, value=0.7, step=0.05, label="Temperature")
189
  btn = gr.Button("Generate", variant="primary")
190
  with gr.Column():
@@ -192,4 +192,10 @@ with gr.Blocks(title="Gemma 2B API") as demo:
192
  btn.click(fn=gradio_generate, inputs=[prompt_box, sys_box, max_tok, temp], outputs=output)
193
 
194
 
 
 
195
  app = gr.mount_gradio_app(api, demo, path="/")
 
 
 
 
 
11
  logging.basicConfig(level=logging.INFO, format="%(asctime)s | %(levelname)s | %(message)s")
12
  logger = logging.getLogger(__name__)
13
 
 
14
  MODEL_ID = os.getenv("MODEL_ID", "facebook/opt-350m")
15
  HF_TOKEN = os.getenv("HF_TOKEN", None)
16
  DEVICE = "cpu"
 
22
  model = AutoModelForCausalLM.from_pretrained(
23
  MODEL_ID,
24
  torch_dtype=DTYPE,
25
+ device_map="cpu",
26
  token=HF_TOKEN,
27
  )
28
  model.eval()
29
  logger.info("Model ready.")
30
 
 
 
 
 
 
 
 
 
 
31
 
32
  class GenerateRequest(BaseModel):
33
  prompt: str = Field(..., min_length=1, max_length=4096)
34
+ max_new_tokens: int = Field(default=128, ge=1, le=512)
35
  temperature: float = Field(default=0.7, ge=0.01, le=2.0)
36
  top_p: float = Field(default=0.9, ge=0.0, le=1.0)
37
  top_k: int = Field(default=50, ge=0, le=200)
 
42
  def stream_tokens(req: GenerateRequest) -> Iterator[str]:
43
  try:
44
  if req.system_prompt:
45
+ prompt = f"{req.system_prompt}\n\nHuman: {req.prompt}\nAssistant:"
 
 
 
 
46
  else:
47
+ prompt = f"Human: {req.prompt}\nAssistant:"
 
 
 
48
 
49
  inputs = tokenizer(
50
+ prompt, return_tensors="pt", truncation=True, max_length=1024
51
+ ).to(DEVICE)
52
 
53
  streamer = TextIteratorStreamer(
54
+ tokenizer, skip_prompt=True, skip_special_tokens=True, timeout=120.0
55
  )
56
 
57
  gen_kwargs = dict(
 
80
 
81
  t.join()
82
  latency = (time.perf_counter() - start) * 1000
83
+ logger.info(f"Done: {token_count} tokens in {latency:.0f}ms")
84
+ yield f"data: {json.dumps({'done': True, 'total_tokens': token_count, 'latency_ms': round(latency, 1)})}\n\n"
85
 
86
  except Exception as e:
87
  tb = traceback.format_exc()
 
89
  yield f"data: {json.dumps({'error': str(e), 'traceback': tb})}\n\n"
90
 
91
 
92
+ # ── FastAPI app ──
93
+ api = FastAPI(title="OPT-350M API")
94
+ api.add_middleware(
95
+ CORSMiddleware,
96
+ allow_origins=["*"],
97
+ allow_credentials=False,
98
+ allow_methods=["*"],
99
+ allow_headers=["*"],
100
+ expose_headers=["*"],
101
+ )
102
+
103
+
104
+ @api.get("/api/health")
105
  async def health():
106
  return {
107
+ "status" : "ok",
108
+ "model" : MODEL_ID,
109
+ "device" : DEVICE,
110
+ "dtype" : str(DTYPE),
 
111
  }
112
 
113
+
114
+ @api.post("/api/generate/stream")
115
  async def generate_stream(req: GenerateRequest):
116
  return StreamingResponse(
117
  stream_tokens(req),
 
123
  },
124
  )
125
 
126
+
127
+ @api.post("/api/generate")
128
  async def generate(req: GenerateRequest):
129
  full_text = ""
130
  total_tokens = 0
131
  latency_ms = 0.0
132
+
133
  for chunk in stream_tokens(req):
134
  if not chunk.startswith("data: "):
135
  continue
 
144
  total_tokens += 1
145
  if "done" in data:
146
  latency_ms = data["latency_ms"]
147
+
148
  return {
149
  "generated_text" : full_text,
150
  "completion_tokens" : total_tokens,
 
153
  }
154
 
155
 
156
+ # ── Gradio UI ──
157
  def gradio_generate(prompt, system_prompt, max_new_tokens, temperature):
158
  req = GenerateRequest(
159
  prompt = prompt,
 
173
  pass
174
 
175
 
176
+ with gr.Blocks(title="OPT-350M API") as demo:
177
  gr.Markdown(
178
+ "## OPT-350M β€” Streaming API\n"
179
  "Use `/api/generate/stream` or `/api/generate` from your backend.\n\n"
180
  "**Health check:** `/api/health`"
181
  )
 
184
  sys_box = gr.Textbox(label="System prompt (optional)", lines=2)
185
  prompt_box = gr.Textbox(label="Prompt", lines=4, placeholder="Ask something...")
186
  with gr.Row():
187
+ max_tok = gr.Slider(32, 512, value=128, step=32, label="Max tokens")
188
  temp = gr.Slider(0.01, 2.0, value=0.7, step=0.05, label="Temperature")
189
  btn = gr.Button("Generate", variant="primary")
190
  with gr.Column():
 
192
  btn.click(fn=gradio_generate, inputs=[prompt_box, sys_box, max_tok, temp], outputs=output)
193
 
194
 
195
+ # ── Mount FastAPI routes into Gradio and launch ──
196
+ # This is the correct way to keep the app alive on HF Spaces
197
  app = gr.mount_gradio_app(api, demo, path="/")
198
+
199
+ if __name__ == "__main__":
200
+ import uvicorn
201
+ uvicorn.run(app, host="0.0.0.0", port=7860)