EvalHarness/evalharness/model/truncation.py
sora d83c2cc1df Serialize tokenizer first load; one-shot degradation warning
96 worker threads racing transformers 5.x lazy imports on the FIRST
_get_tokenizer call raised ImportError and degraded that whole first
batch to the char approximation (the old single-threaded path never
raced). First load now holds a threading.Lock; the transformers
'>model_max_length' logging is silenced inside truncation (counting a
2M-token doc before trimming it is the point), and the per-sample
degradation print becomes a one-shot warning.

Co-Authored-By: Claude <noreply@anthropic.com>
2026-09-15 02:49:34 +00:00

83 lines
3.4 KiB
Python

"""Token-level middle truncation (ported from your /data1/sora/evalscope/bash/run.py).
Keeps head+tail halves of the token stream -- the industry-standard
middle-truncation for long-context benchmarks (longbench_v2 / mrcr).
The evalside run.py uses the same algorithm, guaranteeing comparable inputs.
Usage in run_eval: max_input_tokens=131072 (0/off = no truncation)
Requires a tokenizer (transformers) at tokenizer_path or auto from the model.
"""
import os
import threading
from functools import lru_cache
from typing import Optional
DEFAULT_TRUNCATION_TOKENS = 32768 * 4 # 131072, mirrors evalside run.py
_TOK_LOCK = threading.Lock()
@lru_cache(maxsize=4)
def _get_tokenizer(tokenizer_path: str):
if not tokenizer_path or not os.path.exists(tokenizer_path):
raise FileNotFoundError(
f'tokenizer not found at {tokenizer_path!r} -- token-level truncation '
'needs a local tokenizer dir (e.g. /data1/models/DeepSeek-V4-Flash-INT8)')
# serialize the FIRST load: 96 worker threads racing transformers 5.x's
# lazy imports raised ImportError and silently degraded batches to the
# char approximation; after one success lru_cache serves the rest
with _TOK_LOCK:
from transformers import AutoTokenizer
# the '> model_max_length' warnings are EXPECTED here -- counting a
# 2M-token doc before trimming it is the whole point of truncation
import logging
logging.getLogger('transformers').setLevel(logging.ERROR)
return AutoTokenizer.from_pretrained(tokenizer_path, trust_remote_code=True)
def truncate_middle_tokens(text: str, max_tokens: int, tokenizer_path: str) -> str:
"""Keep head+tail halves of the token stream; decode back to text."""
if max_tokens <= 0 or not text:
return text
tok = _get_tokenizer(tokenizer_path)
ids = tok.encode(text, add_special_tokens=False)
if len(ids) <= max_tokens:
return text
keep_head = max_tokens // 2
keep_tail = max_tokens - keep_head
return tok.decode(ids[:keep_head] + ids[-keep_tail:], skip_special_tokens=True)
def truncate_messages_middle(messages: list, max_tokens: int, tokenizer_path: str,
desired_index: int = 0, window: int = 2) -> list:
"""MRCR-style: when a message list exceeds the budget, keep first/last
messages plus a window around the desired (needle) message, dropping
middle spans; middle of each KEPT long message is token-truncated."""
if max_tokens <= 0 or not messages:
return messages
tok = _get_tokenizer(tokenizer_path)
total = sum(len(tok.encode(m.get('content', '') if isinstance(m, dict) else str(m),
add_special_tokens=False)) for m in messages)
if total <= max_tokens:
return messages
n = len(messages)
keep = set(range(min(2, n))) | set(range(max(0, n - 2), n))
di = desired_index if isinstance(desired_index, int) and 0 <= desired_index < n else 0
keep |= set(range(max(0, di - window), min(n, di + window + 1)))
return [messages[i] for i in sorted(keep)]
def default_tokenizer_path() -> Optional[str]:
"""Candidate local tokenizer for truncation (mirrors evalside default)."""
env = os.environ.get('EVALHARNESS_TOKENIZER')
if env and os.path.exists(env):
return env
for cand in ('/data1/models/DeepSeek-V4-Flash-INT8',):
if os.path.exists(cand):
return cand
return None