refactor generate code
This commit is contained in:
@@ -9,7 +9,10 @@ use crate::{
|
||||
models::common::{InferenceModel, MultiModalData},
|
||||
params::chat::{ChatCompletionChunkResponse, ChatCompletionResponse},
|
||||
tokenizer::TokenizerModel,
|
||||
utils::{build_completion_chunk_response, build_completion_response_with_time},
|
||||
utils::response_utils::{
|
||||
build_chunk_response_with_reasoning, build_chunk_response_with_usage,
|
||||
build_completion_chunk_response, build_completion_response_with_time,
|
||||
},
|
||||
};
|
||||
pub fn get_logit_processor(
|
||||
temperature: Option<f32>,
|
||||
@@ -98,6 +101,17 @@ fn sample_and_push(
|
||||
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,
|
||||
@@ -155,6 +169,7 @@ pub fn generate_stream_generic<M: InferenceModel>(
|
||||
top_k: Option<usize>,
|
||||
seed: u64,
|
||||
max_tokens: u32,
|
||||
in_reasoning: bool,
|
||||
device: &Device,
|
||||
model_name: &str,
|
||||
) -> Result<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
|
||||
@@ -167,14 +182,23 @@ pub fn generate_stream_generic<M: InferenceModel>(
|
||||
max_tokens,
|
||||
device.clone(),
|
||||
);
|
||||
let prompt_tokens = ctx.seq_len as u32;
|
||||
let mut prompt_secs = 0.0f64;
|
||||
let mut completion_tokens = 0u32;
|
||||
let mut completion_secs = 0.0f64;
|
||||
let mut error_tokens = Vec::new();
|
||||
let eos_ids = model.stop_token_ids();
|
||||
let stream = stream! {
|
||||
let mut input_ids = input_ids;
|
||||
let mut tool_call_id = None;
|
||||
let mut tool_call_content = String::new();
|
||||
let mut in_reasoning = in_reasoning;
|
||||
// 处理 unicode 错误累积
|
||||
for _ in 0..ctx.sample_len {
|
||||
let i_start = Instant::now();
|
||||
let logits = if ctx.seqlen_offset == 0 {
|
||||
model.forward_initial(&input_ids, ctx.seqlen_offset, data.clone())
|
||||
|
||||
} else {
|
||||
model.forward_step(&input_ids, ctx.seqlen_offset)
|
||||
}?;
|
||||
@@ -183,6 +207,13 @@ pub fn generate_stream_generic<M: InferenceModel>(
|
||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||
ctx.logit_processor.sample(&logits)?
|
||||
};
|
||||
completion_tokens += 1;
|
||||
let i_duration = i_start.elapsed();
|
||||
if ctx.seqlen_offset == 0 {
|
||||
prompt_secs += i_duration.as_secs_f64();
|
||||
} else {
|
||||
completion_secs += i_duration.as_secs_f64();
|
||||
};
|
||||
|
||||
// 解码(处理�的累积)
|
||||
let decode_ids = if error_tokens.is_empty() {
|
||||
@@ -204,9 +235,60 @@ pub fn generate_stream_generic<M: InferenceModel>(
|
||||
continue;
|
||||
}
|
||||
error_tokens.clear();
|
||||
yield Ok(build_completion_chunk_response(decoded, model_name, None, None));
|
||||
if decoded.eq("<think>") {
|
||||
in_reasoning = true;
|
||||
input_ids = ctx.prepare_for_next_token(next_token)?;
|
||||
continue;
|
||||
}
|
||||
if decoded.eq("</think>") {
|
||||
in_reasoning = false;
|
||||
input_ids = ctx.prepare_for_next_token(next_token)?;
|
||||
continue;
|
||||
}
|
||||
|
||||
// 处理特殊标记和工具调用
|
||||
match decoded.as_str() {
|
||||
"<tool_call>" => {
|
||||
// 开始工具调用
|
||||
tool_call_id = Some(uuid::Uuid::new_v4().to_string());
|
||||
input_ids = ctx.prepare_for_next_token(next_token)?;
|
||||
continue;
|
||||
}
|
||||
"</tool_call>" => {
|
||||
// 结束工具调用
|
||||
let chunk = build_completion_chunk_response(
|
||||
decoded,
|
||||
model_name,
|
||||
tool_call_id.clone(),
|
||||
Some(tool_call_content.clone())
|
||||
);
|
||||
tool_call_id = None;
|
||||
tool_call_content = String::new();
|
||||
yield Ok(chunk);
|
||||
}
|
||||
_ => {
|
||||
if tool_call_id.is_some() {
|
||||
// 在工具调用过程中,收集工具调用内容
|
||||
tool_call_content.push_str(&decoded);
|
||||
input_ids = ctx.prepare_for_next_token(next_token)?;
|
||||
continue;
|
||||
} else {
|
||||
|
||||
// 正常文本输出
|
||||
let chunk = if in_reasoning {
|
||||
build_chunk_response_with_reasoning(decoded, model_name)
|
||||
} else {
|
||||
build_completion_chunk_response(
|
||||
decoded, model_name,
|
||||
None,
|
||||
None
|
||||
)};
|
||||
yield Ok(chunk);
|
||||
}
|
||||
}
|
||||
}
|
||||
if eos_ids.contains(&next_token) {
|
||||
yield Ok(build_chunk_response_with_usage(model_name, completion_tokens.into(), completion_secs.into(), prompt_tokens.into(), prompt_secs.into()));
|
||||
break;
|
||||
}
|
||||
input_ids = ctx.prepare_for_next_token(next_token)?;
|
||||
|
||||
Reference in New Issue
Block a user