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