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}