seng-beliefs / handler.py
nmysore's picture
Upload handler.py with huggingface_hub
1ef6905 verified
Raw History Blame Contribute Delete
8.01 kB
"""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}