File size: 5,921 Bytes
508c74e 4dd4055 a3016ba 4dd4055 a3016ba 24a15a9 4dd4055 8b4bd0d 4dd4055 8b4bd0d 508c74e 8bd5532 508c74e 1c276a8 508c74e 1c276a8 8bd5532 10a1482 1c276a8 8b4bd0d d42cb9c 903aadc 508c74e d42cb9c 4dd4055 160b1b2 508c74e 4dd4055 d42cb9c 4dd4055 160b1b2 4dd4055 d42cb9c 10a1482 d42cb9c 8b4bd0d 24a15a9 8b4bd0d 160b1b2 4dd4055 d42cb9c 4dd4055 a3016ba 4dd4055 d42cb9c 160b1b2 10a1482 4dd4055 8740c4d d42cb9c 160b1b2 b81392c 4dd4055 8b4bd0d 10a1482 4dd4055 8b4bd0d 4dd4055 b81392c 10a1482 | 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 | """
handler.py β HuggingFace Inference Endpoint handler for SriRamanaAtmic/AtmicIntelv1
Compatible with transformers==4.51.3 (matches model's transformers_version in config.json).
Generation parameters (from Section A4 of technical review β do not change):
do_sample = False (greedy decoding β matches SFT + DPO training exactly)
max_new_tokens = 350
repetition_penalty = 1.01 (sole repetition control)
no_repeat_ngram_size = 0 (PERMANENTLY DISABLED β hard-coded, not overridable via API)
temperature / top_p (REMOVED β inactive under greedy decoding)
Token IDs (from added_tokens.json β verified):
<|endoftext|> = 32000 (pad_token_id)
<|assistant|> = 32001 (appears in INPUT prompt β must NEVER be eos_token_id)
<|end|> = 32007 (turn terminator β correct eos for generation)
Critical: generation_config.json in the repo contains eos_token_id=[32000, 32001, 32007].
Token 32001 (<|assistant|>) is present in every input prompt, causing generation to stop
at token 0. This handler explicitly overrides generation_config.json by setting
self.model.generation_config before any generate() call.
Input contract:
The caller (pipeline.py via prompt_builder.py) sends a fully-formatted Phi-3 prompt string.
This handler does NOT apply any chat template β prompt arrives ready to tokenize.
{"inputs": "<|system|>...<|end|>\n<|user|>...<|end|>\n<|assistant|>\n"}
"""
# ββ DynamicCache compatibility shim (transformers >= 4.38) ββββββββββββββββββ
# Must be first β before any other transformers import.
import transformers.cache_utils as _cu
if not hasattr(_cu.DynamicCache, "get_max_length"):
_cu.DynamicCache.get_max_length = lambda self: None
from transformers import DynamicCache
if not hasattr(DynamicCache, "get_max_length"):
DynamicCache.get_max_length = lambda self: None
# ββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
from transformers import AutoTokenizer, AutoModelForCausalLM, GenerationConfig
import torch
class EndpointHandler:
def __init__(self, path=""):
# ββ Tokenizer ββββββββββββββββββββββββββββββββββββββββββββββββββββ
self.tokenizer = AutoTokenizer.from_pretrained(
path,
trust_remote_code=True,
)
# ββ Model ββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
self.model = AutoModelForCausalLM.from_pretrained(
path,
torch_dtype=torch.bfloat16, # matches config.json torch_dtype
device_map="auto",
trust_remote_code=True,
attn_implementation="eager", # avoids flash-attn dependency
)
self.model.eval()
# ββ Override generation_config.json ββββββββββββββββββββββββββββββ
# generation_config.json in the repo has eos_token_id=[32000, 32001, 32007].
# Token 32001 is <|assistant|>, which appears in every input prompt.
# This causes generate() to stop at token 0 β empty output.
# We override it here so model.generate() never reads the repo file.
self.model.generation_config = GenerationConfig(
do_sample=False, # greedy β matches SFT+DPO training
repetition_penalty=1.01,
no_repeat_ngram_size=0, # permanently disabled
eos_token_id=32007, # <|end|> only β turn terminator
pad_token_id=32000, # <|endoftext|>
bos_token_id=1,
)
def __call__(self, data: dict) -> list:
# ββ Input: fully-formatted prompt string from prompt_builder.py ββ
inputs = data.get("inputs", "")
parameters = data.get("parameters", {})
max_new_tokens = int(parameters.get("max_new_tokens", 350))
repetition_penalty = float(parameters.get("repetition_penalty", 1.15))
# ββ Tokenize β prompt already contains all special tokens βββββββββ
tokenized = self.tokenizer(
inputs,
return_tensors="pt",
truncation=True,
max_length=3500, # leaves 596-token headroom within 4096
add_special_tokens=False, # prompt_builder adds
).to(self.model.device)
input_length = tokenized["input_ids"].shape[1]
# ββ Generate ββββββββββββββββββββββββββββββββββββββββββββββββββββββ
# generation_config on the model is already overridden in __init__.
# kwargs here take final precedence for per-request overrides.
with torch.inference_mode():
output = self.model.generate(
**tokenized,
max_new_tokens=max_new_tokens,
repetition_penalty=repetition_penalty,
do_sample=False,
no_repeat_ngram_size=0,
eos_token_id=32007, # <|end|> β confirmed turn terminator
pad_token_id=32000,
)
# ββ Decode new tokens only ββββββββββββββββββββββββββββββββββββββββ
new_tokens = output[0][input_length:]
generated_text = self.tokenizer.decode(new_tokens, skip_special_tokens=True)
return [{"generated_text": generated_text}] |