"""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}