* 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>
49 lines
2.1 KiB
Python
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()])
|