seng-api / server.py
nmysore's picture
Upload folder using huggingface_hub
5b2616c verified
Raw History Blame Contribute Delete
3.17 kB
"""FastAPI server for SENG inference.
Loads the merged GPT-2 Large model (LoRA already folded in) and exposes
endpoints that replicate the exact inference logic from run_pipeline.py.
"""
import os
import torch
from fastapi import FastAPI
from pydantic import BaseModel
from transformers import AutoModelForCausalLM, AutoTokenizer
MODEL_DIR = os.environ.get("MODEL_DIR", "/model")
INSTRUCTION = "Analyze these repository events and extract development beliefs."
CHUNK_SIZE = 12
app = FastAPI(title="SENG API")
# Globals populated at startup
model = None
tokenizer = None
device = None
class PredictRequest(BaseModel):
text: str
class BatchRequest(BaseModel):
narrative: str
@app.on_event("startup")
def load_model():
global model, tokenizer, device
device = "cuda" if torch.cuda.is_available() else "cpu"
dtype = torch.float16 if device == "cuda" else torch.float32
print(f"Loading model from {MODEL_DIR} on {device} ({dtype})")
tokenizer = AutoTokenizer.from_pretrained(MODEL_DIR)
model = AutoModelForCausalLM.from_pretrained(
MODEL_DIR, torch_dtype=dtype
).to(device)
model.eval()
print("Model loaded.")
@app.get("/health")
def health():
return {
"status": "healthy" if model is not None else "loading",
"model_loaded": model is not None,
"gpu_available": torch.cuda.is_available(),
}
def run_inference(text: str) -> str:
"""Run inference on a single text chunk.
Replicates run_pipeline.py:120-146 exactly.
"""
prompt = (
f"### Instruction:\n{INSTRUCTION}\n\n"
f"### Input:\n{text}\n\n"
f"### Response:\n"
)
inputs = tokenizer(
prompt,
return_tensors="pt",
truncation=True,
max_length=768,
).to(device)
with torch.no_grad():
outputs = model.generate(
**inputs,
max_new_tokens=256,
temperature=0.7,
do_sample=True,
top_p=0.9,
pad_token_id=tokenizer.eos_token_id,
repetition_penalty=1.2,
)
generated = tokenizer.decode(
outputs[0][inputs["input_ids"].shape[1]:],
skip_special_tokens=True,
).strip()
return generated
def chunk_narrative(narrative: str, chunk_size: int = CHUNK_SIZE) -> list[str]:
"""Split narrative into non-overlapping chunks.
Replicates run_pipeline.py:89-99.
"""
lines = [line.rstrip("\n") for line in narrative.splitlines() if line.strip()]
chunks = []
for i in range(0, len(lines), chunk_size):
chunk_lines = lines[i : i + chunk_size]
if chunk_lines:
chunks.append("\n".join(chunk_lines))
return chunks
@app.post("/predict")
def predict(req: PredictRequest):
generated = run_inference(req.text)
return {"generated_text": generated}
@app.post("/batch")
def batch(req: BatchRequest):
chunks = chunk_narrative(req.narrative)
results = []
for i, chunk in enumerate(chunks):
generated = run_inference(chunk)
results.append({"chunk_index": i, "generated_text": generated})
return {"total_chunks": len(chunks), "results": results}