Qwen3-4B-AMQ3-Math-SFT / vllm_think_format.py
jepetolee's picture
Add think-format logits processor (tag grammar + budget cut + forced seal)
ce73e24 verified
Raw
History Blame Contribute Delete
13.2 kB
"""think ํƒœ๊ทธ ๋ฌธ๋ฒ• ๊ฐ•์ œ vLLM V1 ๋กœ์ง“ ํ”„๋กœ์„ธ์„œ (2026-07-16).
๊ทผ๊ฑฐ (temp0.75 1Kร—64 ์ „์ˆ˜ ์‹ค์ธก): ์ƒ์„ฑ์˜ 40.1%๊ฐ€ think๋ฅผ ์ œ๋Œ€๋กœ ๋ชป ์—ด๊ณ (34.3%๋Š”
</think>๋ถ€ํ„ฐ ์‹œ์ž‘), ์ •์ƒ ์‹œ์ž‘์กฐ์ฐจ ํƒœ๊ทธ๋ฅผ ํ‰๊ท  2.6๊ฐœ ์‚ฌ์šฉ(์žฌ๊ฐœ๋ฐฉยท์œ ์‚ฌ ๋ฉ€ํ‹ฐํ„ด). ํƒœ๊ทธ
๋ฐฉํ–ฅ ์˜๋ฏธ๋ก ์ด ํ•™์Šต๋˜์ง€ ์•Š์•„ ์„ฑ๋Šฅ๊ณผ ๋ฌด๊ด€ํ•˜๊ฒŒ ํ˜•์‹์ด ๋ถ•๊ดดํ•จ โ€” RL์ด ์ด๋ฅผ ๊ทธ๋Œ€๋กœ ๊ฐ•ํ™”
ํ•˜๊ธฐ ์ „์— ๋ฌธ๋ฒ•์„ ์ƒ์„ฑ ๋‹จ๊ณ„์—์„œ ๊ฐ•์ œํ•œ๋‹ค.
๊ทœ์น™ (ํ† ํฐ id ์ƒํƒœ๋จธ์‹  โ€” ๋””์ฝ”๋“œ ๋ถˆํ•„์š”):
1) think๊ฐ€ ์—ด๋ ค ์žˆ์œผ๋ฉด <think> ์žฌํ˜ธ์ถœ ๊ธˆ์ง€ (์ค‘์ฒฉ/์žฌ๊ฐœ๋ฐฉ ๋ฐฉ์ง€)
2) </think>๊ฐ€ 1ํšŒ ๋“ฑ์žฅํ•œ ์ˆœ๊ฐ„๋ถ€ํ„ฐ <think>ยท</think> ๋ชจ๋‘ ์˜๊ตฌ ๊ธˆ์ง€
3) (์˜ต์…˜) <|im_start|> ๊ธˆ์ง€ โ€” ์ƒˆ ํ„ด ํ™˜๊ฐ ์ฐจ๋‹จ
<think>\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": <int>, # <think> ํ† ํฐ id
"think_close_id": <int>, # </think> ํ† ํฐ id
"prefilled_open": true, # ํ”„๋กฌํ”„ํŠธ๊ฐ€ <think>๋กœ ๋๋‚˜๋Š” ๊ฒฝ์šฐ (ํ”„๋ฆฌํ•„)
"ban_im_start": true, # <|im_start|> ์žฌํ˜ธ์ถœ ๊ธˆ์ง€ (์˜ต์…˜)
"im_start_id": <int>,
}}
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: ํ”„๋กฌํ”„ํŠธ(ํ”„๋ฆฌํ”ฝ์Šค)์— ์ด๋ฏธ <think>โ€ฆ</think>๊ฐ€ ์™„๊ฒฐ๋˜์–ด ์žˆ๋Š” ๊ฒฝ์šฐ
# (์˜ˆ: lrs/MCMC ์ฒญํฌ ์žฌ๊ฐœ โ€” ๋ˆ„์  ํ…์ŠคํŠธ๋ฅผ ํ”„๋กฌํ”„ํŠธ๋กœ ๋„˜๊ธฐ๋Š” ํ›„์† ์š”์ฒญ). ์ด๊ฑธ ์•ˆ ์ฃผ๋ฉด
# ์š”์ฒญ๋งˆ๋‹ค ์ƒํƒœ๊ฐ€ ๋ฆฌ์…‹๋˜์–ด ๋‹ซํžŒ ๋’ค์—๋„ </think> ์žฌ๋ฐฉ์ถœ์ด ํ—ˆ์šฉ๋œ๋‹ค (2026-07-21 ์‹ค์ธก).
self.opened = bool(prefilled_open) or bool(prefilled_closed)
self.closed = bool(prefilled_closed)
self.consumed = 0
# max_think_tokens: <think> ์•ˆ์—์„œ ์ด ํ† ํฐ ์ˆ˜๋ฅผ ๋„˜์œผ๋ฉด ๊ฐ•์ œ ๋ด‰ํ•ฉ
# (non-convergent ์ถ”๋ก  ๋ฃจํ”„ ํƒˆ์ถœ โ†’ ๋‹ต๋ณ€ ๋‹จ๊ณ„๋กœ ๋ฐ€์–ด๋ƒ„). None์ด๋ฉด ๋น„ํ™œ์„ฑ.
self.max_think_tokens = int(max_think_tokens) if max_think_tokens else None
self.think_len = 0 # <think> ์—ด๋ฆฐ ๋’ค ์ƒ์„ฑ๋œ ํ† ํฐ ์ˆ˜
# ๊ฐ•์ œ ๋ด‰ํ•ฉ ํ”„๋ฆฌํ•„ ์‹œํ€€์Šค(์˜ˆ: [</think>, "\n\n"])๋ฅผ ์ˆœ์„œ๋Œ€๋กœ ๊ฐ•์ œํ•œ ๋’ค,
# ์ฒซ ๋‹ต๋ณ€ ํ† ํฐ์€ force_whitelist(ํ•™์Šต๋ฐ์ดํ„ฐ top-N ๋‹ต๋ณ€์‹œ์ž‘ ํ† ํฐ)๋กœ๋งŒ ํ—ˆ์šฉ โ†’
# ๋ชจ๋ธ์ด ๊ทธ์ค‘ ์ตœ๊ณ ๋ฅผ ์Šค์Šค๋กœ ๊ณ ๋ฅด๊ฒŒ. ๋‘˜ ๋‹ค ์—†์œผ๋ฉด </think> ํ•˜๋‚˜๋งŒ ๊ฐ•์ œ(๊ตฌ ๋™์ž‘).
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("<think>")
close_id = tokenizer.convert_tokens_to_ids("</think>")
im_start_id = tokenizer.convert_tokens_to_ids("<|im_start|>")
if open_id is None or close_id is None:
return None
# ๊ฐ•์ œ ๋ด‰ํ•ฉ ํ”„๋ฆฌํ•„: </think> + "\n\n" ์‹œํ€€์Šค ํ›„, ๋‹ต๋ณ€ ์ฒซ ํ† ํฐ์„ ํ•™์Šต๋ฐ์ดํ„ฐ
# top-N ๋‹ต๋ณ€์‹œ์ž‘ ํ† ํฐ(ํ™”์ดํŠธ๋ฆฌ์ŠคํŠธ)์œผ๋กœ๋งŒ ํ—ˆ์šฉํ•ด ๋ชจ๋ธ์ด ์Šค์Šค๋กœ ๊ณ ๋ฅด๊ฒŒ ํ•œ๋‹ค.
# (2026-07-21 amq3 clean 292K ์‹ค์ธก: </think> ๋’ค 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,
}
}