Skip to content

Commit 2e12e75

Browse files
fix: implement repetition_penalty as logits processor — make_sampler lacks it
mlx-lm 0.31.2's make_sampler() does not accept repetition_penalty (only temp, top_p, min_p, top_k, xtc_*). The previous fix passed it as a keyword arg, breaking all generation with TypeError. Now implemented as a proper logits processor that tracks recently generated tokens and penalises their logits (divides positive logits, multiplies negative logits by the penalty factor). Uses a sliding context window of 100 tokens. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
1 parent 928f8c7 commit 2e12e75

2 files changed

Lines changed: 68 additions & 16 deletions

File tree

‎scripts/mlx_generate.py‎

Lines changed: 34 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -72,6 +72,28 @@ def processor(tokens, logits):
7272
return processor
7373

7474

75+
def _make_repetition_processor(penalty, context_size=100):
76+
"""Build a logits processor that penalises recently generated tokens."""
77+
import mlx.core as mx
78+
79+
generated_tokens = []
80+
81+
def processor(tokens, logits):
82+
generated_tokens.append(int(tokens[-1]) if tokens.size > 0 else 0)
83+
recent = generated_tokens[-context_size:]
84+
if recent:
85+
ids = mx.array(list(set(recent)), dtype=mx.int32)
86+
penalties = mx.where(
87+
logits[..., ids] > 0,
88+
logits[..., ids] / penalty,
89+
logits[..., ids] * penalty,
90+
)
91+
logits[..., ids] = penalties
92+
return logits
93+
94+
return processor
95+
96+
7597
def main():
7698
req = json.loads(sys.stdin.read())
7799

@@ -118,24 +140,28 @@ def main():
118140
formatted += system_prompt + "\n\n"
119141
formatted += prompt_text
120142

121-
sampler = make_sampler(
122-
temp=temperature,
123-
top_p=top_p,
124-
repetition_penalty=repetition_penalty if repetition_penalty > 1.0 else None,
125-
repetition_context_size=100,
126-
)
143+
sampler = make_sampler(temp=temperature, top_p=top_p)
127144

128145
# Build logits processor for vocabulary bias
129146
logits_processor = build_logits_processor(logit_bias, tokenizer)
130147

148+
# Chain repetition penalty as a logits processor
149+
processors = []
150+
if repetition_penalty > 1.0:
151+
processors.append(
152+
_make_repetition_processor(repetition_penalty, context_size=100)
153+
)
154+
if logits_processor is not None:
155+
processors.append(logits_processor)
156+
131157
t0 = time.time()
132158

133159
gen_kwargs = dict(
134160
max_tokens=max_tokens,
135161
sampler=sampler,
136162
)
137-
if logits_processor is not None:
138-
gen_kwargs["logits_processors"] = [logits_processor]
163+
if processors:
164+
gen_kwargs["logits_processors"] = processors
139165

140166
full_text = ""
141167
last_resp = None

‎scripts/mlx_worker.py‎

Lines changed: 34 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -58,6 +58,28 @@ def processor(tokens, logits):
5858
return processor
5959

6060

61+
def _make_repetition_processor(penalty, context_size=100):
62+
"""Build a logits processor that penalises recently generated tokens."""
63+
import mlx.core as mx
64+
65+
generated_tokens = []
66+
67+
def processor(tokens, logits):
68+
generated_tokens.append(int(tokens[-1]) if tokens.size > 0 else 0)
69+
recent = generated_tokens[-context_size:]
70+
if recent:
71+
ids = mx.array(list(set(recent)), dtype=mx.int32)
72+
penalties = mx.where(
73+
logits[..., ids] > 0,
74+
logits[..., ids] / penalty,
75+
logits[..., ids] * penalty,
76+
)
77+
logits[..., ids] = penalties
78+
return logits
79+
80+
return processor
81+
82+
6183
def main():
6284
import mlx.core as mx
6385
from mlx_lm import load, stream_generate
@@ -137,17 +159,21 @@ def main():
137159
formatted += system_prompt + "\n\n"
138160
formatted += prompt_text
139161

140-
sampler = make_sampler(
141-
temp=temperature,
142-
top_p=top_p,
143-
repetition_penalty=repetition_penalty if repetition_penalty > 1.0 else None,
144-
repetition_context_size=100,
145-
)
162+
sampler = make_sampler(temp=temperature, top_p=top_p)
146163
logits_processor = build_logits_processor(logit_bias, tokenizer)
147164

148-
gen_kwargs = dict(max_tokens=max_tokens, sampler=sampler)
165+
# Chain repetition penalty as a logits processor
166+
processors = []
167+
if repetition_penalty > 1.0:
168+
processors.append(
169+
_make_repetition_processor(repetition_penalty, context_size=100)
170+
)
149171
if logits_processor is not None:
150-
gen_kwargs["logits_processors"] = [logits_processor]
172+
processors.append(logits_processor)
173+
174+
gen_kwargs = dict(max_tokens=max_tokens, sampler=sampler)
175+
if processors:
176+
gen_kwargs["logits_processors"] = processors
151177

152178
t0 = time.time()
153179

0 commit comments

Comments
 (0)