File size: 8,007 Bytes
103f6c2
 
 
 
 
 
 
 
 
 
 
 
 
e3b3b7e
bc8298c
 
 
103f6c2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
bc8298c
 
 
 
1ef6905
 
 
 
 
 
 
103f6c2
bc8298c
 
 
 
 
 
103f6c2
 
 
 
 
 
 
 
 
 
 
 
 
bc8298c
 
103f6c2
bc8298c
 
103f6c2
 
 
 
 
 
 
 
bc8298c
103f6c2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
"""Custom handler for HF Inference Endpoints.

Loads the pre-merged Gemma 2B IT model and exposes the same inference
logic as server.py — Gemma chat template, INSTRUCTION prompt, and
generation parameters.

Supports MC Dropout for Bayesian uncertainty estimation. Since the model
is pre-merged (no LoRA dropout layers), we inject DropoutWrapper modules
around attention projection layers at init time.
"""

import os
import sys

import torch
from transformers import AutoModelForCausalLM, AutoTokenizer

# On HF Inference Endpoints, handler.py and mc_dropout.py are in the model
# repo root. Ensure that directory is on sys.path for all imports.
_handler_dir = os.path.dirname(os.path.abspath(__file__))
if _handler_dir not in sys.path:
    sys.path.insert(0, _handler_dir)

# MC Dropout import — gracefully degrade if unavailable
_MC_AVAILABLE = False
try:
    from mc_dropout import (
        aggregate_beliefs,
        disable_mc_dropout,
        enable_mc_dropout,
        inject_dropout,
    )
    _MC_AVAILABLE = True
except Exception as _mc_err:
    # Try local development path
    try:
        sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
        from mc_dropout import (
            aggregate_beliefs,
            disable_mc_dropout,
            enable_mc_dropout,
            inject_dropout,
        )
        _MC_AVAILABLE = True
    except Exception:
        print(f"WARNING: MC Dropout not available: {_mc_err}")
        aggregate_beliefs = None
        disable_mc_dropout = None
        enable_mc_dropout = None
        inject_dropout = None

INSTRUCTION = """\
Narrative format: {timestamp} {hash8} {refs} {actors} :: {subject} | {sym}{file}+N/-N@{funcs} ...
Symbols: + added, ~ modified, - deleted, > renamed, = copied
Roles: (a+c) author+committer, (a) author, (c) committer
Funcs: @{name1,name2} for modified functions/classes
Merges: {timestamp} {hash8} MERGE {merged}→{main} :: {subject}
Renames: >{old}→{new}(N%)+N/-N  Large commits: ...+Nmore when >50 files

Read the repository events below. Identify recurring patterns in workflow, code ownership, commit discipline, and architecture.

For each pattern found, output:
Belief: <clear statement about a development practice or pattern>
Evidence:
- Commit: <hash(es) that support this belief>
- Files: <file paths involved>
- Pattern: <what pattern was observed and why it matters>
Confidence: high | medium | low"""

CHUNK_SIZE = 4


class EndpointHandler:
    def __init__(self, path):
        self.device = "cuda" if torch.cuda.is_available() else "cpu"
        # Use bfloat16 on Ampere+ GPUs (A10G, A100) for better numerical stability
        # Fall back to float16 on older GPUs (T4, V100)
        if self.device == "cuda":
            capability = torch.cuda.get_device_capability()
            dtype = torch.bfloat16 if capability[0] >= 8 else torch.float16
        else:
            dtype = torch.float32

        self.tokenizer = AutoTokenizer.from_pretrained(path)
        self.model = AutoModelForCausalLM.from_pretrained(
            path, torch_dtype=dtype
        ).to(self.device)
        self.model.eval()

        # Inject dropout wrappers for MC Dropout support on merged model
        self.mc_dropout_available = False
        if _MC_AVAILABLE and inject_dropout is not None:
            mc_dropout_rate = float(os.environ.get("MC_DROPOUT_RATE", "0.1"))
            n_injected = inject_dropout(self.model, dropout_rate=mc_dropout_rate)
            print(f"Injected {n_injected} dropout wrappers (rate={mc_dropout_rate})")
            self.mc_dropout_available = n_injected > 0
        else:
            print("MC Dropout not available — running in standard mode")

    def _run_inference(self, text):
        """Run inference on a single text chunk using Gemma chat template."""
        user_content = f"{INSTRUCTION}\n\n{text}"
        prompt = (
            f"<start_of_turn>user\n"
            f"{user_content}<end_of_turn>\n"
            f"<start_of_turn>model\n"
        )

        inputs = self.tokenizer(
            prompt,
            return_tensors="pt",
            truncation=True,
            max_length=2048,
        ).to(self.device)

        with torch.no_grad():
            outputs = self.model.generate(
                **inputs,
                max_new_tokens=400,
                do_sample=False,
                pad_token_id=self.tokenizer.eos_token_id,
                eos_token_id=[self.tokenizer.eos_token_id, 107],
            )

        generated = self.tokenizer.decode(
            outputs[0][inputs["input_ids"].shape[1] :],
            skip_special_tokens=True,
        ).strip()

        return generated

    @staticmethod
    def _chunk_narrative(narrative, chunk_size=CHUNK_SIZE):
        """Split narrative into non-overlapping chunks."""
        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

    def _run_mc_inference(self, text, n_passes):
        """Run MC Dropout inference: N passes with dropout enabled."""
        user_content = f"{INSTRUCTION}\n\n{text}"
        prompt = (
            f"<start_of_turn>user\n"
            f"{user_content}<end_of_turn>\n"
            f"<start_of_turn>model\n"
        )

        enable_mc_dropout(self.model)
        pass_texts = []
        for _ in range(n_passes):
            inputs = self.tokenizer(
                prompt,
                return_tensors="pt",
                truncation=True,
                max_length=2048,
            ).to(self.device)

            with torch.no_grad():
                outputs = self.model.generate(
                    **inputs,
                    max_new_tokens=400,
                    temperature=0.7,
                    do_sample=True,
                    top_p=0.9,
                    pad_token_id=self.tokenizer.eos_token_id,
                    eos_token_id=[self.tokenizer.eos_token_id, 107],
                )
            generated = self.tokenizer.decode(
                outputs[0][inputs["input_ids"].shape[1]:],
                skip_special_tokens=True,
            ).strip()
            pass_texts.append(generated)

        disable_mc_dropout(self.model)
        return aggregate_beliefs(pass_texts, n_passes)

    def __call__(self, data):
        inputs = data.get("inputs", "")
        parameters = data.get("parameters", {})
        mode = parameters.get("mode", "predict")
        mc_passes = int(parameters.get("mc_passes", 0))

        # Diagnostic: return handler info when requested
        if parameters.get("info"):
            return {
                "handler_version": "2.0-mc",
                "mc_dropout_available": self.mc_dropout_available,
                "device": str(self.device),
            }

        if mode == "batch":
            chunks = self._chunk_narrative(inputs)
            results = []
            for i, chunk in enumerate(chunks):
                if mc_passes > 0 and self.mc_dropout_available:
                    mc_beliefs = self._run_mc_inference(chunk, mc_passes)
                    results.append({
                        "chunk_index": i,
                        "mc_beliefs": mc_beliefs,
                        "mc_passes": mc_passes,
                    })
                else:
                    generated = self._run_inference(chunk)
                    results.append({"chunk_index": i, "generated_text": generated})
            return {"total_chunks": len(chunks), "results": results}

        if mc_passes > 0 and self.mc_dropout_available:
            mc_beliefs = self._run_mc_inference(inputs, mc_passes)
            return {"mc_beliefs": mc_beliefs, "mc_passes": mc_passes}

        generated = self._run_inference(inputs)
        return {"generated_text": generated}