Skip to content

Commit fcd9d03

Browse files
committed
Return f32 from apply_frequency_presence_penalty, don't round back down
Codex review: the function computed the penalty in f32 but cast the result back to the caller's original dtype before returning. For f16/bf16 logits this could round a small penalty away outright (or noticeably distort it) before LogitsProcessor::sample converts back to f32 anyway for softmax. utils::apply_repeat_penalty already avoids this by always returning f32; do the same here for consistency, including in the zero-penalty no-op path so the function's return dtype is uniform regardless of whether a penalty was actually applied. Adds a regression test with a f16 logit large enough (1000.0, ULP ~0.98) that a 0.5 penalty would be rounded away by a cast back to f16.
1 parent 88facc3 commit fcd9d03

2 files changed

Lines changed: 19 additions & 4 deletions

File tree

‎candle-transformers/src/generation/mod.rs‎

Lines changed: 6 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -26,10 +26,9 @@ pub fn apply_frequency_presence_penalty(
2626
presence_penalty: f32,
2727
) -> Result<Tensor> {
2828
if frequency_penalty == 0. && presence_penalty == 0. {
29-
return Ok(logits.clone());
29+
return logits.to_dtype(DType::F32);
3030
}
3131
let device = logits.device();
32-
let dtype = logits.dtype();
3332
let mut logits_v = logits.to_dtype(DType::F32)?.to_vec1::<f32>()?;
3433
let mut counts = std::collections::HashMap::new();
3534
for token_id in context {
@@ -41,7 +40,11 @@ pub fn apply_frequency_presence_penalty(
4140
}
4241
}
4342
let logits_len = logits_v.len();
44-
Tensor::from_vec(logits_v, logits_len, device)?.to_dtype(dtype)
43+
// Stay in f32: rounding the adjusted logits back down to the caller's original dtype
44+
// (e.g. f16/bf16) can erase a small penalty outright before `LogitsProcessor::sample`
45+
// converts back to f32 anyway. Mirrors `utils::apply_repeat_penalty`, which never casts
46+
// back either.
47+
Tensor::from_vec(logits_v, logits_len, device)
4548
}
4649

4750
pub struct LogitsProcessor {

‎candle-transformers/tests/generation_tests.rs‎

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
use candle::{Device, Result, Tensor};
1+
use candle::{DType, Device, Result, Tensor};
22
use candle_transformers::generation::{
33
apply_frequency_presence_penalty, LogitsProcessor, Sampling,
44
};
@@ -147,6 +147,18 @@ fn frequency_presence_penalty_noop_and_out_of_range() -> Result<()> {
147147
Ok(())
148148
}
149149

150+
#[test]
151+
fn frequency_presence_penalty_preserves_precision_for_f16_input() -> Result<()> {
152+
// A f16 logit of 1000.0 has an ULP of ~0.98, so subtracting a small penalty like 0.5 and
153+
// rounding back down to f16 would erase it outright. Staying in f32 (matching
154+
// `utils::apply_repeat_penalty`'s convention) preserves the adjustment exactly.
155+
let logits = Tensor::new(&[1000.0f32], &Device::Cpu)?.to_dtype(DType::F16)?;
156+
let penalized = apply_frequency_presence_penalty(&logits, &[0], 0.5, 0.0)?;
157+
assert_eq!(penalized.dtype(), DType::F32);
158+
assert_eq!(penalized.to_vec1::<f32>()?, [999.5]);
159+
Ok(())
160+
}
161+
150162
#[test]
151163
fn sample_gumbel() -> Result<()> {
152164
let mut logits_process = LogitsProcessor::from_sampling(

0 commit comments

Comments
 (0)