Xunzhuo's picture
Reuse admitted input encodings without changing model outputs
24368ca verified
Raw History Blame
2.07 kB
"""Complete-input 1K product profile around an unchanged native runtime."""
from dataclasses import replace
MAX_INPUT_TOKENS = 1024
class Complete1KCollator:
"""Use the native collator's actual length error, never truncate fields.
Constructible with a tokenizer before any model is loaded. A whole request
is admitted before predict, and every physical batch before tensor assembly.
"""
def __init__(self, native_collator):
self.base = type(native_collator)(native_collator.tokenizer,
max_length=MAX_INPUT_TOKENS,
state_truncation="error")
self.tokenizer = self.base.tokenizer
self.marker, self.pad = self.base.marker, self.base.pad
self.tensor_batches = 0
def tokens(self, text):
return self.base.tokens(text)
def encode(self, row, labeled=False):
encoded = self.base.encode(row, labeled=labeled)
if (encoded["input_tokens"] > MAX_INPUT_TOKENS
or encoded["state_tokens_original"] != encoded["state_tokens_kept"]):
raise ValueError("Native collator violated the complete-input 1K profile")
return encoded
def admit(self, records):
return [self.encode(row, labeled=False) for row in records]
def __call__(self, records, labeled=False, device="cpu"):
records = list(records)
# All encodes finish before delegating any tensor allocation.
for row in records:
self.encode(row, labeled=labeled)
self.tensor_batches += 1
return self.base(records, labeled=labeled, device=device)
def predict_1k(native, records, *, batch_size=8):
"""Admit every complete input, then reuse its encoding within this request.
Native prediction, batch boundaries and outputs are unchanged. Encodings
are released with the synchronous call; no state activations are cached.
"""
from ._request import predict_1k as predict_admitted
return predict_admitted(native, records, batch_size=batch_size)