unsloth/studio/backend/core/inference/presence_penalty.py
Daniel Han 5608081c35
Studio: apply presence_penalty on the safetensors and MLX inference paths (#6923)
* Studio: apply presence_penalty on the safetensors and MLX inference paths

The safetensors and MLX generate paths resolved the inference config and
then dropped presence_penalty before generation, so the same model applied
the configured value under GGUF and 0 under safetensors/MLX. Thread the
already-resolved presence_penalty through the orchestrator command, worker
gen_kwargs, and the safetensors/MLX generate calls, and apply it with a
small logits processor (subtract once per distinct completion token,
prompt excluded, presence not frequency, zero is a no-op, negatives raise).

Backwards compatible: presence_penalty defaults to 0.0 (byte-identical
output when unset) and the GGUF path is unchanged. Also forward min_p on
the legacy /generate/stream route and add the missing min_p field to
GenerateRequest.

* Studio: bound presence_penalty generated ids to valid vocab range on both paths

The presence-penalty logits processors index by generated token ids. The
torch path filtered only the upper bound (seen < vocab_size), so a negative
id would silently wrap to the wrong row; the MLX path had no bound at all,
and MLX out-of-bounds indexing is documented undefined behavior (crash or
memory corruption on Apple Silicon), unlike torch's harmless negative wrap.

Bound generated ids to [0, vocab) consistently on both paths:
- torch: seen[(seen >= 0) & (seen < vocab_size)] (zero-regression safety net;
  real completion tokens are always in range).
- MLX: route out-of-range/negative ids to a discarded scratch slot via
  mx.where and a (vocab + 1)-wide scatter-assign mask, then subtract. MLX has
  no boolean-mask filtering (data-dependent output shape), so this keeps a
  fixed shape, stays on-device, and preserves once-per-distinct-token
  semantics without any torch/numpy dependency.

Add torch tests for out-of-range and negative ids (only in-range distinct
ids penalized, stray ids ignored, no wrong-index wrap) and a bound-documenting
MLX test that runs on the arm64 macOS CI.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
2026-07-06 22:24:47 -07:00

49 lines
2.1 KiB
Python

# SPDX-License-Identifier: AGPL-3.0-only
# Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"""Presence-penalty logits helpers for the safetensors/MLX inference paths.
Kept in a dependency-light leaf module (torch + transformers only, no unsloth /
peft) so the pure logic can be imported and unit-tested without pulling in the
full inference backend. ``core.inference.inference`` re-exports these for the
runtime generate paths.
"""
import torch
def apply_presence_penalty(input_ids, scores, penalty: float, prompt_len: int):
"""OpenAI/llama.cpp presence penalty: subtract ``penalty`` once per distinct
completion token (positions >= prompt_len; prompt excluded, multiplicity
ignored, negatives raise). In place; zero is a no-op."""
if not penalty:
return scores
vocab_size = scores.shape[-1]
for b in range(input_ids.shape[0]):
generated = input_ids[b, prompt_len:]
if generated.numel() == 0:
continue
seen = torch.unique(generated)
# Bound generated ids to the valid range [0, vocab_size). Real completion
# tokens are always in range, so this is a zero-regression safety net that
# drops any stray out-of-range or negative id before indexing (mirrors the
# MLX path's bound). Filtering both ends avoids indexing scores with a
# negative id (which would silently wrap to the wrong row).
seen = seen[(seen >= 0) & (seen < vocab_size)]
if seen.numel():
scores[b, seen] = scores[b, seen] - penalty
return scores
def _make_presence_penalty_processor(penalty: float, prompt_len: int):
"""``LogitsProcessorList`` for ``apply_presence_penalty``; ``None`` at zero penalty (generate call stays byte-identical)."""
if not penalty:
return None
from transformers import LogitsProcessor, LogitsProcessorList
class _PresencePenaltyLogitsProcessor(LogitsProcessor):
@torch.no_grad()
def __call__(self, input_ids, scores):
return apply_presence_penalty(input_ids, scores, penalty, prompt_len)
return LogitsProcessorList([_PresencePenaltyLogitsProcessor()])