Sampling Strategies for Token Generation in PersonaPlex: Greedy, Top-k, and Nucleus Sampling Explained

PersonaPlex implements four distinct sampling strategies—greedy decoding, temperature-scaled multinomial sampling, top-k filtering, and nucleus (top-p) sampling—within the moshi/moshi/utils/sampling.py module, orchestrated by the LMGen class to control randomness and output diversity during autoregressive generation.

The NVIDIA PersonaPlex repository provides a flexible token generation pipeline that supports both deterministic and stochastic decoding methods. Understanding these sampling strategies is essential for controlling the trade-off between coherence and creativity in generated audio and text tokens. This article examines the implementation details, configuration parameters, and practical usage based on the source code analysis of the sampling utilities and language model generator.

Core Sampling Strategies in PersonaPlex

The sampling pipeline evaluates runtime flags to select from four mutually exclusive generation modes. Each strategy manipulates the output logits produced by the transformer to control which token indices are selected as the next input in the autoregressive sequence.

Greedy Decoding

Greedy decoding represents the deterministic baseline. When use_sampling=False or the temp parameter is less than or equal to 0.0, the system selects the token with the highest logit value using torch.argmax.

In moshi/moshi/utils/sampling.py, the sample_token function implements this behavior at lines 124‑125:

next_token = torch.argmax(logits, dim=-1, keepdim=True)

This mode produces identical outputs for identical inputs, maximizing likelihood but potentially sacrificing diversity or naturalness in conversational audio generation.

Temperature-Scaled Multinomial Sampling

Temperature scaling introduces stochasticity by flattening or sharpening the probability distribution before drawing a sample. The implementation divides logits by a temperature value temp before applying softmax, then draws from the resulting categorical distribution using torch.multinomial.

According to the source code in sample_token (lines 115‑117), the calculation follows:

probs = torch.softmax(logits / temp, dim=-1)

The actual categorical draw is handled by the multinomial helper function (lines 36‑70), which efficiently handles batched probability distributions. Lower temperatures (e.g., 0.7) produce more focused, conservative outputs, while higher values (e.g., 1.0) increase randomness.

Top-k Sampling

Top-k sampling restricts the candidate vocabulary to the k most probable tokens, eliminating long-tail noise from the distribution. The sample_top_k function (lines 72‑84) masks out all logits below the k-th largest value, renormalizes the remaining probabilities, and performs a multinomial draw.

This strategy is invoked when top_k > 0 and no top_p value is specified, offering a fixed-size sampling window that prevents selection of extremely rare tokens regardless of their absolute probability mass.

Nucleus (Top-p) Sampling

Nucleus sampling (also known as top-p sampling) dynamically adjusts the vocabulary size based on cumulative probability mass. The algorithm sorts tokens by probability and retains the smallest prefix whose cumulative probability exceeds the threshold p.

Implemented in sample_top_p (lines 86‑103), this method sorts the probability tensor in descending order, computes cumulative sums, and truncates the distribution at the smallest index where the sum exceeds p. The remaining probabilities are renormalized before the final multinomial selection.

Top-p sampling takes precedence over top-k sampling when both parameters are non-zero, providing a more adaptive approach to vocabulary restriction than the fixed cutoff of top-k.

Sampling Selection Logic and Implementation

The central dispatcher sample_token in moshi/moshi/utils/sampling.py implements a clear precedence hierarchy for selecting between these strategies:

  1. If use_sampling is disabled or temp <= 0.0, execute greedy decoding
  2. Otherwise, apply temperature scaling and check for nucleus constraints
  3. If top_p > 0.0, execute nucleus sampling (ignores top_k)
  4. Else if top_k > 0, execute top-k sampling
  5. Otherwise, execute basic temperature-scaled multinomial sampling

This logic ensures that developers can specify overlapping parameters without ambiguity, as nucleus constraints automatically override fixed-size vocabulary limits.

Integration with the Language Model Generator

The LMGen class in moshi/moshi/models/lm.py serves as the high-level interface that forwards logits to the sampling pipeline. Configured during instantiation, LMGen maintains separate temperature and top-k parameters for audio and text streams:

  • temp: Temperature for audio tokens (default 0.8)
  • temp_text: Temperature for text tokens (default 0.7)
  • top_k: Vocabulary limit for audio (default 250)
  • top_k_text: Vocabulary limit for text (default 25)

These defaults are established in the LMGen constructor (lines 46‑55). During the forward pass, process_transformer_output calls sample_token separately for text and audio logits, applying the modality-specific parameters to maintain appropriate coherence levels across different output types.

Code Examples

Direct Use of Sampling Utilities

For low-level control over token selection, import the sampling functions directly from the utilities module:

import torch
from moshi.moshi.utils.sampling import sample_token

# Simulate transformer output logits for a vocabulary size of 10,000

logits = torch.randn(1, 10000)

# Configure for top-k sampling with temperature

tokens = sample_token(
    logits,
    use_sampling=True,
    temp=0.9,
    top_k=50,
    top_p=0.0,
)

print("Selected token index:", tokens.item())

This example bypasses the generator class to call sample_token directly, triggering the sample_top_k branch due to the non-zero top_k parameter.

Sampling via the Language Model Generator

For production inference, configure sampling parameters through the LMGen wrapper:

from moshi.moshi.models.lm import LMGen, LMModel
import torch

# Assume lm_model is a pre-trained LMModel instance

generator = LMGen(
    lm_model=lm_model,
    device="cuda",
    use_sampling=True,
    temp=0.8,          # Audio temperature

    temp_text=0.7,     # Text temperature

    top_k=250,         # Audio top-k

    top_k_text=25,     # Text top-k

)

# Generate one step from an initial text prompt

output, _ = generator.step(
    input_tokens=None,
    moshi_tokens=None,
    text_token=torch.tensor([[42]], dtype=torch.long),
)

print("Generated sequence shape:", output.shape)

The step method handles the autoregressive cache management while delegating the actual token selection to sample_token with the constructor-specified parameters.

Summary

  • Four strategies: PersonaPlex supports greedy decoding, temperature-scaled multinomial, top-k, and nucleus (top-p) sampling through a unified interface.
  • Central dispatcher: The sample_token function in moshi/moshi/utils/sampling.py routes between strategies based on runtime flags, with top-p taking precedence over top-k when both are specified.
  • Dual-stream configuration: The LMGen class maintains separate sampling parameters for audio (temp, top_k) and text (temp_text, top_k_text) streams, with distinct defaults optimized for each modality.
  • Deterministic fallback: Setting use_sampling=False or temp <= 0.0 disables stochastic methods, ensuring reproducible outputs via torch.argmax.

Frequently Asked Questions

What is the difference between top-k and top-p sampling in PersonaPlex?

Top-k sampling restricts generation to a fixed number of highest-probability tokens regardless of their cumulative probability mass, while top-p (nucleus) sampling dynamically includes tokens until their combined probability exceeds the threshold p. According to the implementation in sample_token, if both parameters are non-zero, top-p takes precedence, making nucleus sampling the active strategy.

How do I enable completely deterministic generation in PersonaPlex?

To disable all randomness and enable greedy decoding, set either use_sampling=False or specify a temperature of 0.0 or below in the LMGen constructor or sample_token call. This triggers the fallback branch that executes torch.argmax(logits, dim=-1, keepdim=True) at lines 124‑125 of sampling.py.

Why does PersonaPlex use different default temperatures for audio and text tokens?

The LMGen constructor initializes temp=0.8 for audio tokens and temp_text=0.7 for text tokens to account for the different statistical properties and coherence requirements of speech versus language generation. Audio sequences typically benefit from slightly higher temperature to maintain natural prosodic variation, while text generation uses a lower temperature to preserve semantic coherence.

Where can I modify the sampling behavior without changing the model weights?

All sampling logic resides in moshi/moshi/utils/sampling.py, specifically within the sample_token, sample_top_k, and sample_top_p functions. For application-level configuration, modify the arguments passed to LMGen in moshi/moshi/models/lm.py or adjust the CLI parameters exposed in moshi/moshi/offline.py (--temp, --topk-audio, --topk-text, --greedy) and moshi/moshi/server.py for HTTP-based inference.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →