generate add penalty repeat

This commit is contained in:
jhqxxx
2026-04-03 14:15:10 +08:00
parent bd8eee6520
commit 279480e3d7
26 changed files with 132 additions and 142 deletions
+27 -22
View File
@@ -49,6 +49,8 @@ pub fn get_logit_processor(
pub struct GenerationContext {
pub logit_processor: LogitsProcessor,
pub repeat_penalty: f32,
pub repeat_last_n: usize,
pub seqlen_offset: usize,
pub seq_len: usize,
pub sample_len: u32,
@@ -60,6 +62,8 @@ impl GenerationContext {
temperature: Option<f32>,
top_p: Option<f32>,
top_k: Option<usize>,
repeat_penalty: Option<f32>,
repeat_last_n: Option<usize>,
seed: u64,
initial_seq_len: usize,
max_tokens: u32,
@@ -67,6 +71,8 @@ impl GenerationContext {
) -> Self {
Self {
logit_processor: get_logit_processor(temperature, top_p, top_k, seed),
repeat_penalty: repeat_penalty.unwrap_or(1.0),
repeat_last_n: repeat_last_n.unwrap_or(64),
seqlen_offset: 0,
seq_len: initial_seq_len,
sample_len: max_tokens,
@@ -91,27 +97,27 @@ impl GenerationContext {
/// 采样辅助函数
fn sample_and_push(
processor: &mut LogitsProcessor,
ctx: &mut GenerationContext,
logits: &Tensor,
generated: &mut Vec<u32>,
) -> Result<u32> {
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
let token = processor.sample(&logits)?;
// 重复惩罚
let logits = if ctx.repeat_penalty == 1. || ctx.repeat_last_n == 0 {
logits
} else {
let start_at = generated.len().saturating_sub(ctx.repeat_last_n);
candle_transformers::utils::apply_repeat_penalty(
&logits,
ctx.repeat_penalty,
&generated[start_at..],
)?
};
let token = ctx.logit_processor.sample(&logits)?;
generated.push(token);
Ok(token)
}
// TODO
// let logits = if self.repeat_penalty == 1. {
// logits
// } else {
// let start_at = generate.len().saturating_sub(self.repeat_last_n);
// candle_transformers::utils::apply_repeat_penalty(
// &logits,
// self.repeat_penalty,
// &generate[start_at..],
// )?
// };
pub fn generate_generic<M: InferenceModel>(
model: &mut M,
tokenizer: &TokenizerModel,
@@ -125,7 +131,7 @@ pub fn generate_generic<M: InferenceModel>(
let eos_ids = model.stop_token_ids();
let i_start = Instant::now();
let logits = model.forward_initial(&input_ids, ctx.seqlen_offset, data)?;
let next_token = sample_and_push(&mut ctx.logit_processor, &logits, &mut generated)?;
let next_token = sample_and_push(ctx, &logits, &mut generated)?;
let i_duration = i_start.elapsed();
let prompt_secs = i_duration.as_secs_f64();
let mut input_ids = ctx.prepare_for_next_token(next_token)?;
@@ -134,12 +140,11 @@ pub fn generate_generic<M: InferenceModel>(
let i_start = Instant::now();
for _ in 1..ctx.sample_len {
let logits = model.forward_step(&input_ids, ctx.seqlen_offset)?;
let next_token = sample_and_push(&mut ctx.logit_processor, &logits, &mut generated)?;
let next_token = sample_and_push(ctx, &logits, &mut generated)?;
if eos_ids.contains(&next_token) {
break;
}
input_ids = ctx.prepare_for_next_token(next_token)?;
}
let i_duration = i_start.elapsed();
@@ -167,6 +172,8 @@ pub fn generate_stream_generic<M: InferenceModel>(
temperature: Option<f32>,
top_p: Option<f32>,
top_k: Option<usize>,
repeat_penalty: Option<f32>,
repeat_last_n: Option<usize>,
seed: u64,
max_tokens: u32,
in_reasoning: bool,
@@ -177,6 +184,8 @@ pub fn generate_stream_generic<M: InferenceModel>(
temperature,
top_p,
top_k,
repeat_penalty,
repeat_last_n,
seed,
input_ids.dim(1)?,
max_tokens,
@@ -193,7 +202,7 @@ pub fn generate_stream_generic<M: InferenceModel>(
let mut tool_call_id = None;
let mut tool_call_content = String::new();
let mut in_reasoning = in_reasoning;
// 处理 unicode 错误累积
let mut generated = Vec::new();
for _ in 0..ctx.sample_len {
let i_start = Instant::now();
let logits = if ctx.seqlen_offset == 0 {
@@ -202,11 +211,7 @@ pub fn generate_stream_generic<M: InferenceModel>(
} else {
model.forward_step(&input_ids, ctx.seqlen_offset)
}?;
let next_token = {
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
ctx.logit_processor.sample(&logits)?
};
let next_token = sample_and_push(&mut ctx, &logits, &mut generated)?;
completion_tokens += 1;
let i_duration = i_start.elapsed();
if ctx.seqlen_offset == 0 {