File size: 3,170 Bytes
16cab96 5b2616c 16cab96 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 | """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}
|