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

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +13 -12
app.py CHANGED
@@ -11,10 +11,11 @@ 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_ID = os.getenv("MODEL_ID", "google/gemma-2b-it")
15
- HF_TOKEN = os.getenv("HF_TOKEN")
16
- DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
17
- DTYPE = torch.float16 if DEVICE == "cuda" else torch.float32
 
18
 
19
  logger.info(f"Loading {MODEL_ID} on {DEVICE} ...")
20
 
@@ -28,7 +29,6 @@ model = AutoModelForCausalLM.from_pretrained(
28
  model.eval()
29
  logger.info("Model ready.")
30
 
31
- # ── FastAPI (mounted under /api) ──
32
  api = FastAPI(title="Gemma 2B API")
33
  api.add_middleware(
34
  CORSMiddleware,
@@ -146,11 +146,14 @@ async def generate(req: GenerateRequest):
146
  total_tokens += 1
147
  if "done" in data:
148
  latency_ms = data["latency_ms"]
149
- return {"generated_text": full_text, "completion_tokens": total_tokens,
150
- "latency_ms": latency_ms, "model": MODEL_ID}
 
 
 
 
151
 
152
 
153
- # ── Gradio UI (required to get free T4 GPU) ──
154
  def gradio_generate(prompt, system_prompt, max_new_tokens, temperature):
155
  req = GenerateRequest(
156
  prompt = prompt,
@@ -173,7 +176,7 @@ def gradio_generate(prompt, system_prompt, max_new_tokens, temperature):
173
  with gr.Blocks(title="Gemma 2B API") as demo:
174
  gr.Markdown(
175
  "## Gemma 2B β€” Streaming API\n"
176
- "Use the `/api/generate/stream` or `/api/generate` endpoints from your backend.\n\n"
177
  "**Health check:** `/api/health`"
178
  )
179
  with gr.Row():
@@ -182,13 +185,11 @@ with gr.Blocks(title="Gemma 2B API") as demo:
182
  prompt_box = gr.Textbox(label="Prompt", lines=4, placeholder="Ask something...")
183
  with gr.Row():
184
  max_tok = gr.Slider(32, 1024, value=256, step=32, label="Max tokens")
185
- temp = gr.Slider(0.01, 2.0, value=0.7, step=0.05, label="Temperature")
186
  btn = gr.Button("Generate", variant="primary")
187
  with gr.Column():
188
  output = gr.Textbox(label="Output", lines=12)
189
  btn.click(fn=gradio_generate, inputs=[prompt_box, sys_box, max_tok, temp], outputs=output)
190
 
191
 
192
- # Mount FastAPI at /api β€” DO NOT call demo.launch() here
193
- # HF Spaces runs its own uvicorn server automatically
194
  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 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"
18
+ DTYPE = torch.float32
19
 
20
  logger.info(f"Loading {MODEL_ID} on {DEVICE} ...")
21
 
 
29
  model.eval()
30
  logger.info("Model ready.")
31
 
 
32
  api = FastAPI(title="Gemma 2B API")
33
  api.add_middleware(
34
  CORSMiddleware,
 
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,
152
+ "latency_ms" : latency_ms,
153
+ "model" : MODEL_ID,
154
+ }
155
 
156
 
 
157
  def gradio_generate(prompt, system_prompt, max_new_tokens, temperature):
158
  req = GenerateRequest(
159
  prompt = prompt,
 
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
  )
182
  with gr.Row():
 
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():
191
  output = gr.Textbox(label="Output", lines=12)
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="/")