How the window_penalty Repetition Suppressor Computes Divisors in YuE
The window_penalty function in YuE computes token-specific divisors by counting occurrences in a recent history window, applying the configured penalty only to repeated tokens; when the penalty equals 1.0, every divisor becomes 1.0, effectively disabling repetition suppression.
The window_penalty repetition suppressor is a critical inference-time mechanism in the multimodal-art-projection/YuE music generation pipeline. Located in the core sampling module, this function dynamically scales model logits to discourage token repetition within a configurable temporal window. Understanding its divisor calculation and edge-case behavior—particularly when the penalty is neutral—is essential for controlling generation diversity.
How the Divisor Is Computed
The divisor calculation follows a three-phase process that transforms the raw logits based on generation history.
Scoping the Penalty Window
The function receives recent_ids, a tensor or list representing the slice of previously generated tokens defined by history[-sampling.penalty_window:]. This window determines how far back the model looks for repetition patterns. Only tokens falling within this specific range influence the divisor calculation.
Counting Occurrences and Building the Divisor
For each token in the vocabulary, the implementation counts how many times that token appears within recent_ids. The divisor for any token t is determined by the following rule:
- If the token appears one or more times in the window, the divisor equals the configured
penaltyvalue - If the token does not appear in the window, the divisor equals
1.0
In src/yue2/sampling.py (lines 16-31), this logic is implemented such that the resulting divisor tensor has the same shape as the logits. Tokens with higher repetition counts do not receive additional scaling—the divisor is binary (penalty or 1.0) rather than cumulative.
Applying Logit Scaling
The function divides the raw model logits by the computed divisor tensor. This operation is equivalent to multiplying by the reciprocal:
# Conceptual implementation
divisor = torch.where(counts > 0, penalty, 1.0)
scaled_logits = logits / divisor
By downscaling logits for recently seen tokens while leaving novel tokens untouched, the function reduces the probability mass assigned to repetitive candidates before temperature scaling or top-p filtering occurs.
Behavior When the Penalty Equals 1.0
When the repetition_penalty parameter is set to 1.0, the divisor for every token—regardless of occurrence count—becomes 1.0. Dividing logits by 1.0 produces no change, meaning the repetition suppressor is effectively disabled. The model treats repeated tokens identically to novel tokens, allowing unrestricted repetition patterns.
This behavior is intentionally permitted by the validation logic in src/yue2/protocol.py (around line 38), which only enforces that the penalty be strictly positive:
if self.repetition_penalty <= 0 or not 1 <= self.penalty_window <= 100:
raise ValueError("Invalid repetition penalty/window")
A value of 1.0 passes this check and results in a no-op penalty application.
Implementation Details in YuE
Source Location and Function Signature
The core implementation resides in src/yue2/sampling.py. The window_penalty function accepts the raw logits, the recent token IDs within the window, and the scalar penalty value. The implementation uses vectorized PyTorch operations to compute the divisor tensor efficiently across the entire vocabulary.
Integration with the Sampling Loop
Within the generation loop (around line 39 in sampling.py), window_penalty is invoked immediately after the model produces raw logits but before temperature scaling and top-p nucleus filtering. This ordering ensures that repetition penalties are applied to the original model outputs, with the modified logits then proceeding through the standard probability distribution shaping.
Practical Code Examples
Basic Usage with Active Penalty
from yue2.sampling import Sampling, window_penalty
import torch
# Configure sampling with repetition suppression
sampling = Sampling(
temperature=0.8,
top_p=0.95,
repetition_penalty=1.2, # Active penalty
penalty_window=64,
)
# Inside the generation loop
logits = model.predict(next_input) # Shape: [batch_size, vocab_size]
# Apply window penalty
logits = window_penalty(
logits,
recent_ids=generated_ids[-sampling.penalty_window:],
penalty=sampling.repetition_penalty,
)
# Proceed with temperature scaling
Disabling Repetition Suppression
# Configure neutral penalty
sampling = Sampling(
temperature=0.9,
top_p=0.9,
repetition_penalty=1.0, # Divisors will all be 1.0
penalty_window=32,
)
# This call leaves logits unchanged
logits = window_penalty(
logits,
recent_ids=generated_ids[-sampling.penalty_window:],
penalty=1.0
)
Debugging Divisor Calculation
import torch
def debug_divisor(logits, recent_ids, penalty):
"""Inspect the divisor tensor for specific tokens."""
# Count occurrences in the window
counts = torch.bincount(
torch.tensor(recent_ids),
minlength=logits.size(-1)
)
# Build divisor: penalty if count>0, else 1.0
divisor = torch.where(
counts > 0,
torch.tensor(penalty, dtype=logits.dtype),
torch.tensor(1.0, dtype=logits.dtype)
)
return divisor
# Example inspection
recent_tokens = [42, 108, 42, 15] # Token 42 appears twice
penalty = 1.5
div = debug_divisor(torch.randn(1000), recent_tokens, penalty)
print(f"Token 42 divisor: {div[42]}") # 1.5
print(f"Token 99 divisor: {div[99]}") # 1.0
Summary
- The divisor equals the configured
penaltyfor tokens present in the recent window, and1.0for all others. - In
src/yue2/sampling.py, thewindow_penaltyfunction applies this divisor by dividing raw logits, reducing probabilities for repeated tokens. - When the penalty is set to
1.0, all divisors become1.0, leaving logits unchanged and effectively disabling the suppressor. - Validation in
src/yue2/protocol.pypermits1.0as a valid neutral value, requiring only that the penalty be positive.
Frequently Asked Questions
Where is the window_penalty function implemented in YuE?
The window_penalty function is implemented in src/yue2/sampling.py (lines 16-31) within the multimodal-art-projection/YuE repository. This module contains the core sampling logic including temperature scaling, top-p filtering, and repetition penalty application.
Why does a penalty of 1.0 disable the repetition suppressor?
A penalty of 1.0 disables suppression because the divisor for every token becomes 1.0. Since the scaling operation divides logits by this divisor, dividing by 1.0 produces no change to the original values. This mathematical identity means recent tokens receive no probability reduction, allowing unrestricted repetition.
How does the penalty window size affect the divisor calculation?
The penalty_window parameter determines how many recent tokens are examined for occurrence counting. A larger window increases the likelihood that any given token will be found in the history, causing more tokens to receive the penalty divisor rather than 1.0. However, the divisor value itself remains constant (the configured penalty) regardless of how many times a token appears within that window.
Is the divisor applied before or after temperature scaling in the sampling pipeline?
The divisor is applied before temperature scaling and top-p filtering. In src/yue2/sampling.py (around line 39), window_penalty processes the raw model logits immediately after generation, with the scaled logits then passing through temperature division and nucleus sampling. This ordering ensures repetition penalties affect the base probability distribution before uncertainty and diversity controls are applied.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →