"""think 태그 문법 강제 vLLM V1 로짓 프로세서 (2026-07-16).
근거 (temp0.75 1K×64 전수 실측): 생성의 40.1%가 think를 제대로 못 열고(34.3%는
부터 시작), 정상 시작조차 태그를 평균 2.6개 사용(재개방·유사 멀티턴). 태그
방향 의미론이 학습되지 않아 성능과 무관하게 형식이 붕괴함 — RL이 이를 그대로 강화
하기 전에 문법을 생성 단계에서 강제한다.
규칙 (토큰 id 상태머신 — 디코드 불필요):
1) think가 열려 있으면 재호출 금지 (중첩/재개방 방지)
2) 가 1회 등장한 순간부터 · 모두 영구 금지
3) (옵션) <|im_start|> 금지 — 새 턴 환각 차단
\n 프리필(ralo.custom_prompts.official_chat_think_prefill_prompt_fn)과 결합 시
문법이 완전 폐쇄된다: 열림 1회(프리필 보장) + 닫힘 정확히 1회 + 이후 답변부.
사용법 (boxed_eos/repetition_abort와 병행 등록 가능):
1) 엔진: vllm_kwargs.logits_processors: ["ralo.vllm_think_format:ThinkFormatLogitsProcessor"]
2) 요청: SamplingParams.extra_args = {"think_format": {
"think_open_id": , # 토큰 id
"think_close_id": , # 토큰 id
"prefilled_open": true, # 프롬프트가 로 끝나는 경우 (프리필)
"ban_im_start": true, # <|im_start|> 재호출 금지 (옵션)
"im_start_id": ,
}}
extra_args에 think_format이 없는 요청은 완전히 무시된다.
주의 — async scheduling 비호환 (2026-07-16 실측): vLLM V1은 async scheduling을 기본
자동 활성화하는데, 이때 워커가 output_tok_ids에 실제 토큰 대신 -1 플레이스홀더를
채운다(gpu_model_runner의 use_async_scheduling 경로). 출력 토큰 값을 읽는 커스텀
프로세서(이 파일 + vllm_boxed_eos + vllm_repetition_abort)는 전부 무력화된다.
반드시 엔진에 `async_scheduling=False`를 함께 줘야 한다. 플레이스홀더가 감지되면
아래 상태머신이 1회 경고를 남긴다.
"""
import logging
from typing import Optional
import torch
try:
from vllm.sampling_params import SamplingParams
from vllm.v1.sample.logits_processor import BatchUpdate, LogitsProcessor
from vllm.v1.sample.logits_processor.builtin import process_dict_updates
_VLLM_OK = True
except ImportError:
_VLLM_OK = False
LogitsProcessor = object # type: ignore
BatchUpdate = None # type: ignore
logger = logging.getLogger(__name__)
_warned_placeholder = False
class ThinkFormatState:
"""토큰 id만으로 금지 목록을 결정하는 상태머신 (HF/vLLM 공용 코어)."""
__slots__ = ("open_id", "close_id", "im_start_id", "opened", "closed", "consumed",
"out", "max_think_tokens", "think_len", "force_prefill",
"force_whitelist", "_force_start")
def __init__(self, out, open_id, close_id, prefilled_open=False, im_start_id=None,
prefilled_closed=False, max_think_tokens=None,
force_prefill=None, force_whitelist=None):
self.out = out # 생성 토큰 리스트 (vLLM 라이브 참조 / HF에선 수동 feed)
self.open_id = int(open_id)
self.close_id = int(close_id)
self.im_start_id = int(im_start_id) if im_start_id is not None else None
# prefilled_closed: 프롬프트(프리픽스)에 이미 …가 완결되어 있는 경우
# (예: lrs/MCMC 청크 재개 — 누적 텍스트를 프롬프트로 넘기는 후속 요청). 이걸 안 주면
# 요청마다 상태가 리셋되어 닫힌 뒤에도 재방출이 허용된다 (2026-07-21 실측).
self.opened = bool(prefilled_open) or bool(prefilled_closed)
self.closed = bool(prefilled_closed)
self.consumed = 0
# max_think_tokens: 안에서 이 토큰 수를 넘으면 강제 봉합
# (non-convergent 추론 루프 탈출 → 답변 단계로 밀어냄). None이면 비활성.
self.max_think_tokens = int(max_think_tokens) if max_think_tokens else None
self.think_len = 0 # 열린 뒤 생성된 토큰 수
# 강제 봉합 프리필 시퀀스(예: [, "\n\n"])를 순서대로 강제한 뒤,
# 첫 답변 토큰은 force_whitelist(학습데이터 top-N 답변시작 토큰)로만 허용 →
# 모델이 그중 최고를 스스로 고르게. 둘 다 없으면 하나만 강제(구 동작).
self.force_prefill = [int(x) for x in force_prefill] if force_prefill else [self.close_id]
self.force_whitelist = [int(x) for x in force_whitelist] if force_whitelist else None
self._force_start = None # 강제 봉합 시작 시점의 out 길이
def advance(self):
"""새 토큰 소비 후 현재 시점의 금지 토큰 id 리스트 반환."""
global _warned_placeholder
while self.consumed < len(self.out):
t = self.out[self.consumed]
self.consumed += 1
if t == -1 and not _warned_placeholder:
_warned_placeholder = True
logger.warning(
"[ThinkFormat] output_tok_ids에 -1 플레이스홀더 감지 — vLLM async "
"scheduling이 켜져 있어 닫힘 감지가 불가능합니다. 엔진에 "
"async_scheduling=False를 전달하세요.")
if self.opened and not self.closed:
self.think_len += 1
if t == self.open_id:
self.opened = True
elif t == self.close_id and self.opened:
self.closed = True
return self.banned_ids()
def force_allowed_ids(self):
"""강제 봉합 진행 중이면 이번 스텝 허용 토큰 id 리스트, 아니면 None.
프리필 시퀀스를 순서대로 1개씩 강제 → 끝나면 답변 첫 토큰을 화이트리스트로
1스텝 제한 → 그 뒤 해제(None). 진입 후 self.closed가 True가 돼도 _force_start
기준으로 계속 진행한다."""
if self.max_think_tokens is None:
return None
if self._force_start is None:
if (self.opened and not self.closed
and self.think_len >= self.max_think_tokens):
self._force_start = len(self.out) # 강제 봉합 개시
else:
return None
progress = len(self.out) - self._force_start
seq = self.force_prefill
if progress < len(seq):
return [seq[progress]] # 프리필: 그 자리 토큰만 허용
if self.force_whitelist and progress == len(seq):
return list(self.force_whitelist) # 답변 첫 토큰: top-N만 허용
return None # 프리필+화이트리스트 끝 → 해제
def banned_ids(self):
banned = []
if self.closed:
banned = [self.open_id, self.close_id] # 닫힌 후엔 둘 다 영구 금지
elif self.opened:
banned = [self.open_id] # 열려 있는 동안 재개방 금지
if self.im_start_id is not None:
banned.append(self.im_start_id)
return banned
class ThinkFormatLogitsProcessor(LogitsProcessor):
"""think 태그 문법 강제 — 금지 토큰 로짓만 -inf, 그 외 무변경."""
def __init__(self, vllm_config, device: torch.device, is_pin_memory: bool):
if not _VLLM_OK:
raise RuntimeError("vLLM V1 logits processor API를 찾을 수 없음")
self.device = device
self.pin_memory = is_pin_memory
self.states: dict[int, ThinkFormatState] = {}
self._rows: list[int] = []
self._cols: list[int] = []
self._rows_t = None
self._cols_t = None
# 강제 봉합: force_rows는 전체 -inf, (allow_rows, allow_cols)만 0으로 살림
self._force_rows_t = None
self._allow_rows_t = None
self._allow_cols_t = None
def is_argmax_invariant(self) -> bool:
return False
@staticmethod
def add_request(params: "SamplingParams", _prompt, output_tok_ids) -> Optional[ThinkFormatState]:
cfg = (getattr(params, "extra_args", None) or {}).get("think_format")
if not cfg or cfg.get("think_open_id") is None or cfg.get("think_close_id") is None:
return None
return ThinkFormatState(
out=output_tok_ids,
open_id=cfg["think_open_id"],
close_id=cfg["think_close_id"],
prefilled_open=bool(cfg.get("prefilled_open", False)),
prefilled_closed=bool(cfg.get("prefilled_closed", False)),
im_start_id=cfg.get("im_start_id") if cfg.get("ban_im_start", False) else None,
max_think_tokens=cfg.get("max_think_tokens"),
force_prefill=cfg.get("force_prefill"),
force_whitelist=cfg.get("force_whitelist"),
)
def update_state(self, batch_update: "BatchUpdate | None") -> None:
process_dict_updates(self.states, batch_update, self.add_request)
rows, cols = [], [] # 일반 금지(-inf)
force_rows, allow_rows, allow_cols = [], [], [] # 강제: 행 전체 -inf 후 allow만 0
for idx, st in self.states.items():
banned = st.advance()
allowed = st.force_allowed_ids()
if allowed is not None:
force_rows.append(idx)
for tid in allowed:
allow_rows.append(idx)
allow_cols.append(tid)
else:
for tid in banned:
rows.append(idx)
cols.append(tid)
def _t(vals):
return torch.tensor(vals, device="cpu", dtype=torch.int64,
pin_memory=self.pin_memory).to(self.device, non_blocking=True)
self._rows_t, self._cols_t = (_t(rows), _t(cols)) if rows else (None, None)
if force_rows:
self._force_rows_t = _t(force_rows)
self._allow_rows_t = _t(allow_rows)
self._allow_cols_t = _t(allow_cols)
else:
self._force_rows_t = self._allow_rows_t = self._allow_cols_t = None
def apply(self, logits: torch.Tensor) -> torch.Tensor:
if self._rows_t is not None:
logits[self._rows_t, self._cols_t] = float("-inf")
if self._force_rows_t is not None:
# 강제 봉합: 해당 행 전체 -inf 후 허용 토큰만 0 (프리필=1개, 화이트리스트=N개)
logits[self._force_rows_t] = float("-inf")
logits[self._allow_rows_t, self._allow_cols_t] = 0.0
return logits
def build_think_format_extra_args(algo_cfg: dict, tokenizer,
prefilled_open: bool = False) -> Optional[dict]:
"""dapo_kwargs/lrs_kwargs의 think_format 블록 → SamplingParams.extra_args."""
cfg = (algo_cfg or {}).get("think_format") or {}
if not cfg.get("enabled", False):
return None
open_id = tokenizer.convert_tokens_to_ids("")
close_id = tokenizer.convert_tokens_to_ids("")
im_start_id = tokenizer.convert_tokens_to_ids("<|im_start|>")
if open_id is None or close_id is None:
return None
# 강제 봉합 프리필: + "\n\n" 시퀀스 후, 답변 첫 토큰을 학습데이터
# top-N 답변시작 토큰(화이트리스트)으로만 허용해 모델이 스스로 고르게 한다.
# (2026-07-21 amq3 clean 292K 실측: 뒤 100% "\n\n", 첫 실토큰 top20이 97%)
force_prefill = force_whitelist = None
if cfg.get("force_close_prefill"):
nn = tokenizer("\n\n", add_special_tokens=False)["input_ids"]
force_prefill = [int(close_id)] + [int(x) for x in nn]
# config에서 직접 준 화이트리스트(id 리스트) 우선, 없으면 학습데이터 기반 기본값.
# 사람이름(John/Mary/James/Maria)·단일문자는 제외 — 답변 요약이 아니라 스토리
# 문제 재서술로 빠질 수 있어서. 일반 답변시작 토큰만 (커버리지 ~96%):
# To/We/Let/The/Given/###/(/First/In
force_whitelist = cfg.get("force_close_whitelist") or [
1249, 1654, 10061, 785, 22043, 14374, 7, 5338, 641,
]
force_whitelist = [int(x) for x in force_whitelist]
return {
"think_format": {
"think_open_id": int(open_id),
"think_close_id": int(close_id),
"prefilled_open": bool(cfg.get("prefilled_open", prefilled_open)),
"prefilled_closed": bool(cfg.get("prefilled_closed", False)),
"ban_im_start": bool(cfg.get("ban_im_start", True)),
"im_start_id": int(im_start_id) if im_start_id is not None else None,
"max_think_tokens": (int(cfg["max_think_tokens"])
if cfg.get("max_think_tokens") else None),
"force_prefill": force_prefill,
"force_whitelist": force_whitelist,
}
}