"""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)