generate add penalty repeat
This commit is contained in:
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user