-
Notifications
You must be signed in to change notification settings - Fork 7
Expand file tree
/
Copy pathdecoding.py
More file actions
89 lines (76 loc) · 3.48 KB
/
Copy pathdecoding.py
File metadata and controls
89 lines (76 loc) · 3.48 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
"""Decoding configuration with the ``enable_thinking=False`` invariant.
Every decode call in StudentSim sets
``enable_thinking=False`` on Qwen3's chat template. Without that, Qwen3's chat
template defaults to opening a ``<think>`` block that often does not close
within ``max_new_tokens``, so the model never emits the answer token. The
invariant lives on this dataclass so it cannot accidentally be dropped by a
caller; the only way to bypass it is to pass ``allow_thinking=True`` explicitly
(used by no production code path).
"""
from __future__ import annotations
from dataclasses import dataclass
@dataclass(frozen=True, slots=True)
class DecodingConfig:
"""Per-call decoding parameters.
The defaults decode greedily, with the thinking block turned off. Each
domain's own values live in :class:`studentsim.eval.protocol.EvalProtocol`,
which is where evaluation reads them from.
Parameters
----------
max_new_tokens
Token budget. Chess and math use ``32`` (short categorical responses); L2
uses ``256`` (free-form essay fragments).
temperature
Sampling temperature. ``0.0`` is greedy and is the default everywhere in
evaluation; tutor RL rollouts override to ``1.0``.
top_p
Nucleus sampling threshold. Ignored when ``do_sample=False``.
repetition_penalty
Per-token repetition penalty. L2 uses ``1.1`` because pure greedy decoding
on free-form essays occasionally entered repeat loops past the natural
essay end on a non-trivial fraction of outputs; ``1.0`` everywhere else.
do_sample
``False`` means greedy; ``True`` enables sampling via ``temperature`` /
``top_p``.
enable_thinking
ALWAYS ``False`` in production. See module docstring.
"""
max_new_tokens: int
temperature: float = 0.0
top_p: float = 1.0
repetition_penalty: float = 1.0
do_sample: bool = False
enable_thinking: bool = False
def __post_init__(self) -> None:
if self.max_new_tokens <= 0:
raise ValueError(f"max_new_tokens must be positive, got {self.max_new_tokens}")
if not (0.0 <= self.temperature):
raise ValueError(f"temperature must be non-negative, got {self.temperature}")
if not (0.0 < self.top_p <= 1.0):
raise ValueError(f"top_p must be in (0, 1], got {self.top_p}")
if self.repetition_penalty <= 0:
raise ValueError(f"repetition_penalty must be positive, got {self.repetition_penalty}")
def as_hf_kwargs(self) -> dict:
"""Render as ``transformers.GenerationConfig`` keyword arguments."""
kwargs: dict = {
"max_new_tokens": self.max_new_tokens,
"do_sample": self.do_sample,
"repetition_penalty": self.repetition_penalty,
}
if self.do_sample:
kwargs["temperature"] = self.temperature
kwargs["top_p"] = self.top_p
return kwargs
def as_chat_template_kwargs(self) -> dict:
"""Render as ``tokenizer.apply_chat_template`` keyword arguments."""
return {"enable_thinking": self.enable_thinking}
def as_vllm_kwargs(self) -> dict:
"""Render as vLLM ``SamplingParams`` keyword arguments.
vLLM's default uses greedy when ``temperature=0`` regardless of ``do_sample``.
"""
return {
"max_tokens": self.max_new_tokens,
"temperature": self.temperature,
"top_p": self.top_p,
"repetition_penalty": self.repetition_penalty,
}