use ferrum_interfaces::sampler::*;
use ferrum_types::{SamplingParams, TokenId};
use rand::RngCore;
#[test]
fn temperature_processor_scales_logits() {
let mut logits = vec![1.0, 2.0, 3.0];
let params = SamplingParams {
temperature: 2.0,
..Default::default()
};
let prev: Vec<TokenId> = vec![];
let freqs = std::collections::HashMap::new();
let vocab_size = 3;
let mut ctx = SamplingContext::new(0, ¶ms, &mut logits, &prev, &freqs, vocab_size);
let p = TemperatureProcessor::new(2.0);
p.process(&mut ctx).unwrap();
assert!((ctx.logits[2] - 1.5).abs() < 1e-6);
}
#[test]
fn topk_masks_tail() {
let mut logits = vec![0.0, 1.0, 2.0, 3.0];
let params = SamplingParams::default();
let prev = vec![];
let freqs = std::collections::HashMap::new();
let mut ctx = SamplingContext::new(0, ¶ms, &mut logits, &prev, &freqs, 4);
TopKProcessor::new(2).process(&mut ctx).unwrap();
let masked = ctx
.logits
.iter()
.filter(|v| **v == f32::NEG_INFINITY)
.count();
assert!(masked >= 2);
}
fn apply_top_k(mut logits: Vec<f32>, k: usize) -> Vec<f32> {
let params = SamplingParams::default();
let frequencies = std::collections::HashMap::new();
let vocab_size = logits.len();
let mut context = SamplingContext::new(0, ¶ms, &mut logits, &[], &frequencies, vocab_size);
TopKProcessor::new(k).process(&mut context).unwrap();
logits
}
fn logit_bits(logits: &[f32]) -> Vec<u32> {
logits.iter().map(|value| value.to_bits()).collect()
}
fn legacy_top_k(logits: &mut [f32], k: usize) {
if k == 0 || k >= logits.len() {
return;
}
let mut indices = (0..logits.len()).collect::<Vec<_>>();
indices.sort_by(|&left, &right| {
logits[right]
.partial_cmp(&logits[left])
.unwrap_or(std::cmp::Ordering::Equal)
});
let threshold = logits[indices[k - 1]];
for logit in logits {
if *logit < threshold {
*logit = f32::NEG_INFINITY;
}
}
}
#[test]
fn topk_keeps_all_threshold_ties_in_original_token_positions() {
let logits = vec![2.0, -3.0, 9.0, 2.0, 1.0, 2.0];
let filtered = apply_top_k(logits, 2);
assert_eq!(
filtered,
[2.0, f32::NEG_INFINITY, 9.0, 2.0, f32::NEG_INFINITY, 2.0]
);
}
#[test]
fn topk_preserves_signed_zero_bits_and_infinite_thresholds() {
let logits = vec![f32::NEG_INFINITY, -0.0, f32::INFINITY, 0.0, -2.0];
assert_eq!(
logit_bits(&apply_top_k(logits.clone(), 2)),
logit_bits(&[
f32::NEG_INFINITY,
-0.0,
f32::INFINITY,
0.0,
f32::NEG_INFINITY
])
);
assert_eq!(
apply_top_k(
vec![f32::INFINITY, 5.0, f32::INFINITY, f32::NEG_INFINITY],
1
),
[
f32::INFINITY,
f32::NEG_INFINITY,
f32::INFINITY,
f32::NEG_INFINITY
]
);
assert_eq!(
apply_top_k(vec![f32::NEG_INFINITY; 4], 2),
[f32::NEG_INFINITY; 4]
);
}
#[test]
fn topk_disabled_or_out_of_range_preserves_every_bit() {
let logits = vec![f32::from_bits(0x7fc0_1234), -0.0, 0.0, f32::INFINITY, -5.0];
for k in [0, logits.len(), usize::MAX] {
assert_eq!(
logit_bits(&apply_top_k(logits.clone(), k)),
logit_bits(&logits)
);
}
assert!(apply_top_k(Vec::new(), 0).is_empty());
assert!(apply_top_k(Vec::new(), 20).is_empty());
}
#[test]
fn topk_retention_matches_order_statistics_for_finite_logits() {
let logits = (0..97)
.map(|i| ((i * 37) % 17) as f32 - 8.0)
.collect::<Vec<_>>();
for k in 1..logits.len() {
let filtered = apply_top_k(logits.clone(), k);
for (index, (&before, &after)) in logits.iter().zip(&filtered).enumerate() {
let greater = logits.iter().filter(|&&value| value > before).count();
let expected = if greater < k {
before
} else {
f32::NEG_INFINITY
};
assert_eq!(after.to_bits(), expected.to_bits(), "k={k}, token={index}");
}
}
}
#[test]
fn topk_nan_inputs_keep_legacy_mask_and_nan_payloads() {
let nan = f32::from_bits(0x7fc0_1234);
let negative_nan = f32::from_bits(0xffc0_5678);
for logits in [
vec![nan, 3.0, 1.0, 2.0],
vec![3.0, nan, 1.0, 2.0],
vec![3.0, 1.0, 2.0, negative_nan],
vec![
nan,
f32::INFINITY,
-0.0,
negative_nan,
0.0,
f32::NEG_INFINITY,
],
] {
for k in 1..logits.len() {
let mut expected = logits.clone();
legacy_top_k(&mut expected, k);
assert_eq!(
logit_bits(&apply_top_k(logits.clone(), k)),
logit_bits(&expected),
"k={k}"
);
}
}
}
#[test]
fn stochastic_sampling_keeps_seeded_draws_and_temperature_topk_topp_order() {
let params = SamplingParams {
temperature: 0.6,
top_p: 0.95,
top_k: Some(20),
repetition_penalty: 1.0,
seed: Some(20_260_912),
..Default::default()
};
let config = SamplingConfig::from_params(¶ms);
assert_eq!(
config.processor_chain.processor_names(),
["temperature", "top_k", "top_p"]
);
let frequencies = std::collections::HashMap::new();
let mut actual_rng = SamplingRng::seeded(params.seed.unwrap());
let mut expected_rng = actual_rng.clone();
for step in 0..32 {
let mut actual = (0..127)
.map(|i| ((i * 37 + step * 11) % 53) as f32 / 8.0)
.collect::<Vec<_>>();
let mut expected = actual.clone();
let vocab_size = actual.len();
let context =
SamplingContext::new(step, ¶ms, &mut actual, &[], &frequencies, vocab_size);
let actual_token = config.sample(context, &mut actual_rng).unwrap();
for value in &mut expected {
*value /= params.temperature;
}
legacy_top_k(&mut expected, 20);
let mut context =
SamplingContext::new(step, ¶ms, &mut expected, &[], &frequencies, vocab_size);
TopPProcessor::new(params.top_p)
.process(&mut context)
.unwrap();
let expected_token = MultinomialSampler
.sample_with_context(&context, &mut expected_rng)
.unwrap();
assert_eq!(logit_bits(&actual), logit_bits(&expected), "step={step}");
assert_eq!(actual_token, expected_token, "step={step}");
}
assert_eq!(
actual_rng.next_u64(),
expected_rng.next_u64(),
"RNG consumption changed"
);
}
#[test]
fn topp_masks_beyond_p() {
let mut logits = vec![0.0, 0.0, 10.0, 9.0];
let params = SamplingParams {
top_p: 0.6,
..Default::default()
};
let binding = std::collections::HashMap::new();
let mut ctx = SamplingContext::new(0, ¶ms, &mut logits, &[], &binding, 4);
TopPProcessor::new(0.6).process(&mut ctx).unwrap();
let masked = ctx
.logits
.iter()
.filter(|v| **v == f32::NEG_INFINITY)
.count();
assert!(masked >= 1);
}
#[test]
fn topp_keeps_the_minimal_prefix_when_mass_equals_threshold() {
let mut logits = vec![0.0; 4];
let params = SamplingParams {
top_p: 0.5,
..Default::default()
};
let frequencies = std::collections::HashMap::new();
let mut ctx = SamplingContext::new(0, ¶ms, &mut logits, &[], &frequencies, 4);
TopPProcessor::new(0.5).process(&mut ctx).unwrap();
assert_eq!(
ctx.logits.iter().filter(|logit| logit.is_finite()).count(),
2
);
}
#[test]
fn c13_captured_logits_have_a_stable_sampling_contract() {
const VOCAB_SIZE: usize = 248_320;
const TOP20: [(usize, f32); 20] = [
(760, 18.859375),
(31_248, 18.0625),
(40, 15.875),
(2_064, 15.46875),
(248_069, 14.671875),
(1_206, 14.59375),
(248_058, 14.5859375),
(31_391, 14.1171875),
(1_919, 14.0390625),
(69_060, 14.0390625),
(96_747, 13.6640625),
(90_700, 13.65625),
(4_272, 13.5859375),
(20_740, 13.546875),
(97_237, 13.515625),
(8_160, 13.375),
(56_757, 13.3046875),
(27, 13.1328125),
(2_381, 12.96875),
(21_765, 12.9296875),
];
let params = SamplingParams {
temperature: 1.0,
top_p: 0.95,
top_k: Some(20),
presence_penalty: 1.5,
repetition_penalty: 1.0,
seed: Some(9_271),
..Default::default()
};
let plan = SamplingConfig::from_params(¶ms);
let mut logits = vec![f32::NEG_INFINITY; VOCAB_SIZE];
for (token_id, logit) in TOP20 {
logits[token_id] = logit;
}
let frequencies = std::collections::HashMap::new();
let mut ctx = SamplingContext::new(0, ¶ms, &mut logits, &[], &frequencies, VOCAB_SIZE);
plan.processor_chain.process(&mut ctx).unwrap();
let finite_token_ids = ctx
.logits
.iter()
.enumerate()
.filter_map(|(token_id, logit)| logit.is_finite().then_some(token_id))
.collect::<Vec<_>>();
assert_eq!(
finite_token_ids,
vec![40, 760, 1_206, 2_064, 31_248, 248_069]
);
let mut rng = SamplingRng::seeded(9_271);
let token = plan.sampler.sample_with_context(&ctx, &mut rng).unwrap();
assert_eq!(token, TokenId::new(31_248));
}
#[test]
fn sampling_rng_stream_is_versioned() {
assert_eq!(
SamplingRng::algorithm_id(),
"chacha12-rand-core-pcg32-u64-v1"
);
let mut rng = SamplingRng::seeded(9_271);
assert_eq!(
[rng.next_u64(), rng.next_u64()],
[8_484_894_384_277_192_859, 5_896_309_912_339_806_864],
"update this vector only with an explicit sampling RNG contract version bump"
);
}
#[test]
fn repetition_penalty_applies() {
let mut logits = vec![1.0, 1.0, -2.0];
let params = SamplingParams::default();
let prev = vec![TokenId::new(1), TokenId::new(1), TokenId::new(2)];
let mut freqs = std::collections::HashMap::new();
freqs.insert(TokenId::new(1), 2usize);
freqs.insert(TokenId::new(2), 1usize);
let vocab_size = 3;
let mut ctx = SamplingContext::new(0, ¶ms, &mut logits, &prev, &freqs, vocab_size);
RepetitionPenaltyProcessor::new(1.1)
.process(&mut ctx)
.unwrap();
assert_eq!(ctx.logits[0], 1.0);
assert!((ctx.logits[1] - (1.0 / 1.1)).abs() < 1e-6);
assert!((ctx.logits[2] - (-2.0 * 1.1)).abs() < 1e-6);
}
#[test]
fn presence_and_frequency_penalties_follow_generated_token_counts() {
let mut logits = vec![1.0, 4.0, -2.0, 8.0];
let params = SamplingParams::default();
let previous = vec![TokenId::new(1), TokenId::new(1), TokenId::new(2)];
let frequencies =
std::collections::HashMap::from([(TokenId::new(1), 2usize), (TokenId::new(2), 1usize)]);
let mut ctx = SamplingContext::new(
previous.len(),
¶ms,
&mut logits,
&previous,
&frequencies,
4,
);
PresenceFrequencyPenaltyProcessor::new(0.5, 0.25)
.process(&mut ctx)
.unwrap();
assert_eq!(ctx.logits[0], 1.0);
assert!((ctx.logits[1] - 3.0).abs() < 1e-6);
assert!((ctx.logits[2] - (-2.75)).abs() < 1e-6);
assert_eq!(ctx.logits[3], 8.0);
}
#[test]
fn negative_presence_and_frequency_penalties_promote_seen_tokens() {
let mut logits = vec![1.0, 2.0];
let params = SamplingParams::default();
let frequencies = std::collections::HashMap::from([(TokenId::new(0), 3usize)]);
let mut ctx = SamplingContext::new(3, ¶ms, &mut logits, &[], &frequencies, 2);
PresenceFrequencyPenaltyProcessor::new(-0.5, -0.25)
.process(&mut ctx)
.unwrap();
assert!((ctx.logits[0] - 2.25).abs() < 1e-6);
assert_eq!(ctx.logits[1], 2.0);
}
#[test]
fn sampling_config_uses_penalty_temperature_and_filter_order() {
let params = SamplingParams {
temperature: 2.0,
top_p: 0.9,
top_k: Some(8),
min_p: Some(0.1),
repetition_penalty: 1.1,
presence_penalty: 0.5,
frequency_penalty: 0.25,
..Default::default()
};
let config = SamplingConfig::from_params(¶ms);
assert_eq!(
config.processor_chain.processor_names(),
vec![
"repetition_penalty",
"presence_frequency_penalty",
"temperature",
"min_p",
"top_k",
"top_p",
]
);
}
#[test]
fn min_p_uses_temperature_scaled_probability_ratio() {
let params = SamplingParams {
temperature: 2.0,
min_p: Some(0.5),
..Default::default()
};
let config = SamplingConfig::from_params(¶ms);
let mut logits = vec![4.0, 3.0, 2.0];
let frequencies = std::collections::HashMap::new();
let mut ctx = SamplingContext::new(0, ¶ms, &mut logits, &[], &frequencies, 3);
config.processor_chain.process(&mut ctx).unwrap();
assert_eq!(ctx.logits[0], 2.0);
assert_eq!(ctx.logits[1], 1.5);
assert_eq!(ctx.logits[2], f32::NEG_INFINITY);
}
#[test]
fn min_p_one_keeps_only_tokens_tied_for_maximum() {
let params = SamplingParams {
min_p: Some(1.0),
..Default::default()
};
let config = SamplingConfig::from_params(¶ms);
let mut logits = vec![3.0, 2.0, 3.0];
let frequencies = std::collections::HashMap::new();
let mut ctx = SamplingContext::new(0, ¶ms, &mut logits, &[], &frequencies, 3);
config.processor_chain.process(&mut ctx).unwrap();
assert_eq!(ctx.logits, &[3.0, f32::NEG_INFINITY, 3.0]);
}
#[test]
fn only_unprocessed_greedy_plan_supports_raw_speculation() {
assert!(
SamplingConfig::from_params(&SamplingParams::greedy()).supports_raw_greedy_speculation()
);
for params in [
SamplingParams {
temperature: 1.0,
..SamplingParams::greedy()
},
SamplingParams {
presence_penalty: 0.1,
..SamplingParams::greedy()
},
SamplingParams {
min_p: Some(0.9),
..SamplingParams::greedy()
},
] {
assert!(!SamplingConfig::from_params(¶ms).supports_raw_greedy_speculation());
}
}
#[test]
fn sampling_config_from_params_builds_chain_and_samples() {
let params = SamplingParams {
temperature: 0.7,
top_p: 0.9,
top_k: Some(10),
..Default::default()
};
let cfg = SamplingConfig::from_params(¶ms);
let mut rng = SamplingRng::seeded(42);
let mut logits = vec![0.1, 0.2, 3.0, 0.4];
let prev: Vec<TokenId> = vec![];
let freqs = std::collections::HashMap::new();
let vocab = logits.len();
let ctx = SamplingContext::new(0, ¶ms, &mut logits, &prev, &freqs, vocab);
let tok = cfg.sample(ctx, &mut rng).unwrap();
assert!((tok.get() as usize) < logits.len());
}
#[test]
fn greedy_sampler_picks_max() {
let g = GreedySampler;
let tok = g
.sample(&[0.1, 10.0, 0.2], &mut SamplingRng::seeded(1))
.unwrap();
assert_eq!(tok, TokenId::new(1));
}
#[test]
fn greedy_sampler_rejects_logits_without_a_finite_candidate() {
let sampler = GreedySampler;
let mut rng = SamplingRng::seeded(1);
for logits in [
vec![f32::NEG_INFINITY; 4],
vec![f32::NAN; 4],
vec![f32::NEG_INFINITY, f32::NAN, f32::INFINITY],
] {
let error = sampler
.sample(&logits, &mut rng)
.expect_err("non-finite logits must not select an arbitrary token");
assert!(
error
.to_string()
.contains("No finite logits available for sampling"),
"{error}"
);
}
}
#[test]
fn greedy_sampler_ignores_non_finite_logits_around_a_valid_candidate() {
let sampler = GreedySampler;
let token = sampler
.sample(
&[f32::NAN, f32::NEG_INFINITY, 3.0, f32::INFINITY],
&mut SamplingRng::seeded(1),
)
.unwrap();
assert_eq!(token, TokenId::new(2));
}