"""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: Evidence: - Commit: - Files: - Pattern: 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"user\n" f"{user_content}\n" f"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"user\n" f"{user_content}\n" f"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}