From 279480e3d7acca596bf0690bef2901d195c41622 Mon Sep 17 00:00:00 2001
From: jhqxxx <18280426169@163.com>
Date: Fri, 3 Apr 2026 14:15:10 +0800
Subject: [PATCH] generate add penalty repeat
---
README.md | 6 +-
README.zh-CN.md | 4 ++
docs/changelog.md | 6 +-
docs/changelog.zh-CN.md | 4 ++
docs/installation.md | 3 +-
docs/installation.zh-CN.md | 3 +-
docs/supported-models.md | 5 +-
docs/supported-models.zh-CN.md | 5 +-
src/cli/mod.rs | 3 +
src/models/common/generate.rs | 49 +++++++------
src/models/common/model_mapping.rs | 3 +
src/models/deepseek_ocr/generate.rs | 4 ++
src/models/fun_asr_nano/generate.rs | 6 +-
src/models/glm_asr_nano/generate.rs | 4 ++
src/models/glm_ocr/generate.rs | 4 ++
src/models/hunyuan_ocr/generate.rs | 4 ++
src/models/lfm2/generate.rs | 4 ++
src/models/lfm2vl/generate.rs | 4 ++
src/models/minicpm4/generate.rs | 4 ++
src/models/mod.rs | 4 ++
src/models/paddleocr_vl/generate.rs | 4 ++
src/models/qwen3/generate.rs | 4 ++
src/models/qwen3_5/generate.rs | 16 +++--
src/models/qwen3vl/generate.rs | 107 ++--------------------------
src/params/chat.rs | 10 +++
tests/test_glm_ocr.rs | 4 +-
26 files changed, 132 insertions(+), 142 deletions(-)
diff --git a/README.md b/README.md
index dc9599a..697379f 100644
--- a/README.md
+++ b/README.md
@@ -46,6 +46,10 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
- **🧠 Attention Optimization** - Optional Flash Attention support for optimized long sequence processing
## Changelog
+### 2026-04-03
+- ChatCompletionParameters add repeat_penalty and repeat_last_n
+- generate add penalty repeat code
+
### 2026-04-02
- refactor generate code
- \...\ The content of the thought chain is returned using the reasoning_content field.
@@ -55,7 +59,7 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
- refactor deepseek_ocr/fun_asr_nano generate code
### 2026-03-31
-- add server adn cli mod
+- add server and cli mod
- aha model name use modelscope id replace
- update WhichModel
- Usage add time info
diff --git a/README.zh-CN.md b/README.zh-CN.md
index 386e946..ef33e2c 100644
--- a/README.zh-CN.md
+++ b/README.zh-CN.md
@@ -46,6 +46,10 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
## 更新日志
## Changelog
+### 2026-04-03
+- ChatCompletionParameters 新增 repeat_penalty 和 repeat_last_n 参数
+- 生成添加重复惩罚代码
+
### 2026-04-02
- 重构生成代码
- \...\ 思维链内容使用reasoning_content字段返回。
diff --git a/docs/changelog.md b/docs/changelog.md
index bc980ed..44f314a 100644
--- a/docs/changelog.md
+++ b/docs/changelog.md
@@ -5,6 +5,10 @@ All notable changes to aha will be documented in this file.
The format is based on [Keep a Changelog](https://keepachangelog.com/en/1.0.0/),
and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0.html).
+### 2026-04-03
+- ChatCompletionParameters add repeat_penalty and repeat_last_n
+- generate add penalty repeat code
+
### 2026-04-02
- refactor generate code
- \...\ The content of the thought chain is returned using the reasoning_content field.
@@ -14,7 +18,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
- refactor deepseek_ocr/fun_asr_nano generate code
### 2026-03-31
-- add server adn cli mod
+- add server and cli mod
- aha model name use modelscope id replace
- update WhichModel enum
- Usage add time info
diff --git a/docs/changelog.zh-CN.md b/docs/changelog.zh-CN.md
index 4d5bf0c..84646f1 100644
--- a/docs/changelog.zh-CN.md
+++ b/docs/changelog.zh-CN.md
@@ -5,6 +5,10 @@
格式基于 [Keep a Changelog](https://keepachangelog.com/zh-CN/1.0.0/),
本项目遵循 [语义化版本](https://semver.org/lang/zh-CN/spec/v2.0.0.html)。
+### 2026-04-03
+- ChatCompletionParameters 新增 repeat_penalty 和 repeat_last_n 参数
+- 生成添加重复惩罚代码
+
### 2026-04-02
- 重构生成代码
- \...\ 思维链内容使用reasoning_content字段返回。
diff --git a/docs/installation.md b/docs/installation.md
index 3c3bca9..55faff4 100644
--- a/docs/installation.md
+++ b/docs/installation.md
@@ -290,7 +290,8 @@ LiquidAI/LFM2-VL-1.6B LiquidAI vlm ✔
OpenBMB/MiniCPM4-0.5B OpenBMB llm ✔
Qwen/Qwen2.5-VL-3B-Instruct Qwen vlm ✔
Qwen/Qwen2.5-VL-7B-Instruct Qwen vlm
-Qwen/Qwen3-0.6B Qwen llm ✔
+Qwen/Qwen3-0.6B Qwen llm ✔
+Qwen/Qwen3-1.7B Qwen llm ✔
Qwen/Qwen3.5-0.8B Qwen vlm ✔
Qwen/Qwen3.5-2B Qwen vlm
Qwen/Qwen3.5-4B Qwen vlm
diff --git a/docs/installation.zh-CN.md b/docs/installation.zh-CN.md
index e7f0aea..7726e41 100644
--- a/docs/installation.zh-CN.md
+++ b/docs/installation.zh-CN.md
@@ -290,7 +290,8 @@ LiquidAI/LFM2-VL-1.6B LiquidAI vlm ✔
OpenBMB/MiniCPM4-0.5B OpenBMB llm ✔
Qwen/Qwen2.5-VL-3B-Instruct Qwen vlm ✔
Qwen/Qwen2.5-VL-7B-Instruct Qwen vlm
-Qwen/Qwen3-0.6B Qwen llm ✔
+Qwen/Qwen3-0.6B Qwen llm ✔
+Qwen/Qwen3-1.7B Qwen llm ✔
Qwen/Qwen3.5-0.8B Qwen vlm ✔
Qwen/Qwen3.5-2B Qwen vlm
Qwen/Qwen3.5-4B Qwen vlm
diff --git a/docs/supported-models.md b/docs/supported-models.md
index 0b4414b..f7a1edb 100644
--- a/docs/supported-models.md
+++ b/docs/supported-models.md
@@ -14,7 +14,8 @@ LiquidAI/LFM2-VL-1.6B LiquidAI vlm ✔
OpenBMB/MiniCPM4-0.5B OpenBMB llm ✔
Qwen/Qwen2.5-VL-3B-Instruct Qwen vlm ✔
Qwen/Qwen2.5-VL-7B-Instruct Qwen vlm
-Qwen/Qwen3-0.6B Qwen llm ✔
+Qwen/Qwen3-0.6B Qwen llm ✔
+Qwen/Qwen3-1.7B Qwen llm ✔
Qwen/Qwen3.5-0.8B Qwen vlm ✔
Qwen/Qwen3.5-2B Qwen vlm
Qwen/Qwen3.5-4B Qwen vlm
@@ -43,7 +44,7 @@ ZhipuAI/GLM-OCR ZhipuAI ocr ✔
| Model | Parameters | Model Id | License |
|-------|-----------|--------|---------|
-| **Qwen3-0.6B** | 0.6B | Qwen/Qwen3-0.6B | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
+| **Qwen3-0.6B** | 0.6B | Qwen/Qwen3-0.6B
Qwen/Qwen3-1.7B | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **MiniCPM4-0.5B** | 0.5B | OpenBMB/MiniCPM4-0.5B | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **LFM2-1.2B** | 1.2B | LiquidAI/LFM2-1.2B | [lfm1.0](https://huggingface.co/LiquidAI/LFM2-1.2B/blob/main/LICENSE) |
| **LFM2.5-1.2B-Instruct** | 1.2B | LiquidAI/LFM2.5-1.2B-Instruct | [lfm1.0](https://huggingface.co/LiquidAI/LFM2.5-1.2B-Instruct/blob/main/LICENSE) |
diff --git a/docs/supported-models.zh-CN.md b/docs/supported-models.zh-CN.md
index b451d02..071df14 100644
--- a/docs/supported-models.zh-CN.md
+++ b/docs/supported-models.zh-CN.md
@@ -14,7 +14,8 @@ LiquidAI/LFM2-VL-1.6B LiquidAI vlm ✔
OpenBMB/MiniCPM4-0.5B OpenBMB llm ✔
Qwen/Qwen2.5-VL-3B-Instruct Qwen vlm ✔
Qwen/Qwen2.5-VL-7B-Instruct Qwen vlm
-Qwen/Qwen3-0.6B Qwen llm ✔
+Qwen/Qwen3-0.6B Qwen llm ✔
+Qwen/Qwen3-1.7B Qwen llm ✔
Qwen/Qwen3.5-0.8B Qwen vlm ✔
Qwen/Qwen3.5-2B Qwen vlm
Qwen/Qwen3.5-4B Qwen vlm
@@ -43,7 +44,7 @@ ZhipuAI/GLM-OCR ZhipuAI ocr ✔
| 模型 | 参数量 | 模型id | 开源协议 |
|------|--------|------|---------|
-| **Qwen3-0.6B** | 0.6B | Qwen/Qwen3-0.6B | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
+| **Qwen3-0.6B** | 0.6B | Qwen/Qwen3-0.6B
Qwen/Qwen3-1.7B | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **MiniCPM4-0.5B** | 0.5B | OpenBMB/MiniCPM4-0.5B | [Apache 2.0](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/apache-2.0.md) |
| **LFM2-1.2B** | 1.2B | LiquidAI/LFM2-1.2B | [lfm1.0](https://huggingface.co/LiquidAI/LFM2-1.2B/blob/main/LICENSE) |
| **LFM2.5-1.2B-Instruct** | 1.2B | LiquidAI/LFM2.5-1.2B-Instruct | [lfm1.0](https://huggingface.co/LiquidAI/LFM2.5-1.2B-Instruct/blob/main/LICENSE) |
diff --git a/src/cli/mod.rs b/src/cli/mod.rs
index c9f0323..12d910d 100644
--- a/src/cli/mod.rs
+++ b/src/cli/mod.rs
@@ -240,6 +240,9 @@ pub(crate) fn run_run(args: RunArgs) -> anyhow::Result<()> {
WhichModel::Qwen3_0_6B => {
qwen3::Qwen3Exec::run(&input, output.as_deref(), &weight_path)?;
}
+ WhichModel::Qwen3_1_7B => {
+ qwen3::Qwen3Exec::run(&input, output.as_deref(), &weight_path)?;
+ }
WhichModel::Qwen3_5_0_8B => {
qwen3_5::Qwen3_5Exec::run(&input, output.as_deref(), &weight_path)?;
}
diff --git a/src/models/common/generate.rs b/src/models/common/generate.rs
index 1f34d8b..b80083f 100644
--- a/src/models/common/generate.rs
+++ b/src/models/common/generate.rs
@@ -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,
top_p: Option,
top_k: Option,
+ repeat_penalty: Option,
+ repeat_last_n: Option,
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,
) -> Result {
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(
model: &mut M,
tokenizer: &TokenizerModel,
@@ -125,7 +131,7 @@ pub fn generate_generic(
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(
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(
temperature: Option,
top_p: Option,
top_k: Option,
+ repeat_penalty: Option,
+ repeat_last_n: Option,
seed: u64,
max_tokens: u32,
in_reasoning: bool,
@@ -177,6 +184,8 @@ pub fn generate_stream_generic(
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(
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(
} 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 {
diff --git a/src/models/common/model_mapping.rs b/src/models/common/model_mapping.rs
index c84da6f..0fa5d99 100644
--- a/src/models/common/model_mapping.rs
+++ b/src/models/common/model_mapping.rs
@@ -18,6 +18,8 @@ pub enum WhichModel {
Qwen2_5VL7B,
#[value(name = "Qwen/Qwen3-0.6B")]
Qwen3_0_6B,
+ #[value(name = "Qwen/Qwen3-1.7B")]
+ Qwen3_1_7B,
#[value(name = "Qwen/Qwen3.5-0.8B")]
Qwen3_5_0_8B,
#[value(name = "Qwen/Qwen3.5-2B")]
@@ -93,6 +95,7 @@ impl WhichModel {
// LLM models
WhichModel::MiniCPM4_0_5B
| WhichModel::Qwen3_0_6B
+ | WhichModel::Qwen3_1_7B
| WhichModel::LFM2_1_2B
| WhichModel::LFM2_5_1_2BInstruct => "llm",
// VLM models
diff --git a/src/models/deepseek_ocr/generate.rs b/src/models/deepseek_ocr/generate.rs
index 939afaf..9427684 100644
--- a/src/models/deepseek_ocr/generate.rs
+++ b/src/models/deepseek_ocr/generate.rs
@@ -98,6 +98,8 @@ impl GenerateModel for DeepseekOCRGenerateModel {
mes.temperature,
mes.top_p,
None,
+ mes.repeat_penalty,
+ mes.repeat_last_n,
mes.seed.unwrap_or(34562) as u64,
input_ids.dim(1)?,
max_tokens,
@@ -169,6 +171,8 @@ impl GenerateModel for DeepseekOCRGenerateModel {
temperature,
top_p,
None,
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
max_tokens,
false,
diff --git a/src/models/fun_asr_nano/generate.rs b/src/models/fun_asr_nano/generate.rs
index 446648d..e7ef42c 100644
--- a/src/models/fun_asr_nano/generate.rs
+++ b/src/models/fun_asr_nano/generate.rs
@@ -109,6 +109,8 @@ impl GenerateModel for FunAsrNanoGenerateModel {
temperature.into(),
top_p.into(),
top_k.into(),
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
input_ids.dim(1)?,
max_tokens,
@@ -145,8 +147,6 @@ impl GenerateModel for FunAsrNanoGenerateModel {
let top_k = self.generation_config.top_k;
let seed = mes.seed.unwrap_or(34562) as u64;
let max_tokens = mes.max_tokens.unwrap_or(1024);
- // let mut logit_processor =
- // get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
let (speech, fbank_mask, input_ids) = self.processor.process_info(&mes, &self.tokenizer)?;
let speech = speech.to_dtype(self.dtype)?;
let data_vec = vec![speech.into(), fbank_mask.into()];
@@ -159,6 +159,8 @@ impl GenerateModel for FunAsrNanoGenerateModel {
temperature.into(),
top_p.into(),
top_k.into(),
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
max_tokens,
false,
diff --git a/src/models/glm_asr_nano/generate.rs b/src/models/glm_asr_nano/generate.rs
index 66eada0..117239d 100644
--- a/src/models/glm_asr_nano/generate.rs
+++ b/src/models/glm_asr_nano/generate.rs
@@ -77,6 +77,8 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
mes.temperature,
mes.top_p,
mes.top_k,
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
input_ids.dim(1)?,
sample_len,
@@ -124,6 +126,8 @@ impl<'a> GenerateModel for GlmAsrNanoGenerateModel<'a> {
mes.temperature,
mes.top_p,
None,
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
sample_len,
false,
diff --git a/src/models/glm_ocr/generate.rs b/src/models/glm_ocr/generate.rs
index e3b3c94..1b3b6df 100644
--- a/src/models/glm_ocr/generate.rs
+++ b/src/models/glm_ocr/generate.rs
@@ -113,6 +113,8 @@ impl GenerateModel for GlmOcrGenerateModel {
mes.temperature,
mes.top_p,
mes.top_k,
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
input_ids.dim(1)?,
sample_len,
@@ -187,6 +189,8 @@ impl GenerateModel for GlmOcrGenerateModel {
mes.temperature,
mes.top_p,
None,
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
sample_len,
false,
diff --git a/src/models/hunyuan_ocr/generate.rs b/src/models/hunyuan_ocr/generate.rs
index 8b0e212..df5d7a4 100644
--- a/src/models/hunyuan_ocr/generate.rs
+++ b/src/models/hunyuan_ocr/generate.rs
@@ -87,6 +87,8 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
temperature.into(),
top_p.into(),
top_k.into(),
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
input_ids.dim(1)?,
sample_len,
@@ -149,6 +151,8 @@ impl<'a> GenerateModel for HunyuanOCRGenerateModel<'a> {
temperature.into(),
top_p.into(),
top_k.into(),
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
sample_len,
false,
diff --git a/src/models/lfm2/generate.rs b/src/models/lfm2/generate.rs
index d81d794..3020a3a 100644
--- a/src/models/lfm2/generate.rs
+++ b/src/models/lfm2/generate.rs
@@ -72,6 +72,8 @@ impl<'a> GenerateModel for Lfm2GenerateModel<'a> {
mes.temperature,
mes.top_p,
None,
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
input_ids.dim(1)?,
sample_len,
@@ -114,6 +116,8 @@ impl<'a> GenerateModel for Lfm2GenerateModel<'a> {
mes.temperature,
mes.top_p,
None,
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
sample_len,
false,
diff --git a/src/models/lfm2vl/generate.rs b/src/models/lfm2vl/generate.rs
index e78c6f3..afa7183 100644
--- a/src/models/lfm2vl/generate.rs
+++ b/src/models/lfm2vl/generate.rs
@@ -72,6 +72,8 @@ impl<'a> GenerateModel for Lfm2VLGenerateModel<'a> {
mes.temperature,
mes.top_p,
None,
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
input_ids.dim(1)?,
sample_len,
@@ -126,6 +128,8 @@ impl<'a> GenerateModel for Lfm2VLGenerateModel<'a> {
mes.temperature,
mes.top_p,
None,
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
sample_len,
false,
diff --git a/src/models/minicpm4/generate.rs b/src/models/minicpm4/generate.rs
index d586c34..98a5758 100644
--- a/src/models/minicpm4/generate.rs
+++ b/src/models/minicpm4/generate.rs
@@ -64,6 +64,8 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
mes.temperature,
mes.top_p,
None,
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
input_ids.dim(1)?,
sample_len,
@@ -104,6 +106,8 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
mes.temperature,
mes.top_p,
None,
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
sample_len,
false,
diff --git a/src/models/mod.rs b/src/models/mod.rs
index ddcaf58..b52fc86 100644
--- a/src/models/mod.rs
+++ b/src/models/mod.rs
@@ -168,6 +168,10 @@ pub fn load_model<'a>(
let model = Qwen3GenerateModel::init(path, None, None)?;
ModelInstance::Qwen3(model)
}
+ WhichModel::Qwen3_1_7B => {
+ let model = Qwen3GenerateModel::init(path, None, None)?;
+ ModelInstance::Qwen3(model)
+ }
WhichModel::Qwen3_5_0_8B => {
let model = Qwen3_5GenerateModel::init(path, None, None)?;
ModelInstance::Qwen3_5(model)
diff --git a/src/models/paddleocr_vl/generate.rs b/src/models/paddleocr_vl/generate.rs
index 5a39a03..0c8cd87 100644
--- a/src/models/paddleocr_vl/generate.rs
+++ b/src/models/paddleocr_vl/generate.rs
@@ -79,6 +79,8 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
mes.temperature,
mes.top_p,
None,
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
input_ids.dim(1)?,
sample_len,
@@ -141,6 +143,8 @@ impl<'a> GenerateModel for PaddleOCRVLGenerateModel<'a> {
mes.temperature,
mes.top_p,
None,
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
sample_len,
false,
diff --git a/src/models/qwen3/generate.rs b/src/models/qwen3/generate.rs
index 5c65191..b0c13e9 100644
--- a/src/models/qwen3/generate.rs
+++ b/src/models/qwen3/generate.rs
@@ -71,6 +71,8 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
temperature.into(),
top_p.into(),
top_k.into(),
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
input_ids.dim(1)?,
sample_len,
@@ -117,6 +119,8 @@ impl<'a> GenerateModel for Qwen3GenerateModel<'a> {
temperature.into(),
top_p.into(),
top_k.into(),
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
sample_len,
in_reasoning,
diff --git a/src/models/qwen3_5/generate.rs b/src/models/qwen3_5/generate.rs
index 25dc8d2..7064719 100644
--- a/src/models/qwen3_5/generate.rs
+++ b/src/models/qwen3_5/generate.rs
@@ -29,8 +29,8 @@ pub struct Qwen3_5GenerateModel<'a> {
qwen3_5: Qwen3_5Model,
device: Device,
model_name: String,
- // repeat_penalty: f32, // TODO
- // repeat_last_n: usize,
+ repeat_penalty: f32,
+ repeat_last_n: usize,
}
impl<'a> Qwen3_5GenerateModel<'a> {
@@ -59,8 +59,8 @@ impl<'a> Qwen3_5GenerateModel<'a> {
qwen3_5,
device,
model_name: model_name.to_string(),
- // repeat_penalty: 1.01,
- // repeat_last_n: 64,
+ repeat_penalty: 1.01,
+ repeat_last_n: 64,
})
}
@@ -116,8 +116,8 @@ impl<'a> Qwen3_5GenerateModel<'a> {
qwen3_5,
device,
model_name: stem.to_string(),
- // repeat_penalty: 1.1,
- // repeat_last_n: 64,
+ repeat_penalty: 1.1,
+ repeat_last_n: 64,
})
}
}
@@ -147,6 +147,8 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
temperature.into(),
top_p.into(),
Some(20),
+ self.repeat_penalty.into(),
+ self.repeat_last_n.into(),
seed,
input_ids.dim(1)?,
sample_len,
@@ -213,6 +215,8 @@ impl<'a> GenerateModel for Qwen3_5GenerateModel<'a> {
mes.temperature,
mes.top_p,
None,
+ self.repeat_penalty.into(),
+ self.repeat_last_n.into(),
seed,
sample_len,
in_reasoning,
diff --git a/src/models/qwen3vl/generate.rs b/src/models/qwen3vl/generate.rs
index f4b91ca..b87dc6a 100644
--- a/src/models/qwen3vl/generate.rs
+++ b/src/models/qwen3vl/generate.rs
@@ -84,6 +84,8 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
temperature.into(),
top_p.into(),
top_k.into(),
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
input_ids.dim(1)?,
sample_len,
@@ -150,6 +152,8 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
temperature.into(),
top_p.into(),
top_k.into(),
+ mes.repeat_penalty,
+ mes.repeat_last_n,
seed,
sample_len,
in_reasoning,
@@ -157,108 +161,5 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
&self.model_name,
)?;
Ok(Box::new(Box::pin(stream)))
- // let sample_len = mes.max_tokens.unwrap_or(1024);
- // let stream = stream! {
- // let mut error_tokens = Vec::new();
- // let mut pixel_values = input.pixel_values.as_ref();
- // let image_grid_thw = input.image_grid_thw.as_ref();
- // let mut pixel_values_video = input.pixel_values_video.as_ref();
- // let video_grid_thw = input.video_grid_thw.as_ref();
- // let mut tool_call_id = None;
- // let mut tool_call_content = String::new();
- // for _ in 0..sample_len {
- // let logits = self.qwen3_vl.forward(
- // &input_ids,
- // pixel_values,
- // image_grid_thw,
- // pixel_values_video,
- // video_grid_thw,
- // Some(&cache_position),
- // seqlen_offset,
- // )?;
- // let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
- // let next_token = logit_processor.sample(&logits)?;
- // let mut decode_ids = Vec::new();
- // if !error_tokens.is_empty() {
- // decode_ids.extend_from_slice(&error_tokens);
- // }
- // decode_ids.push(next_token);
- // let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{e}")))?;
- // if decoded_token.contains("�") {
- // error_tokens.push(next_token);
- // if error_tokens.len() > 3 {
- // error_tokens.clear();
- // }
- // seqlen_offset += seq_len;
- // seq_len = 1;
- // input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
- // cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
- // pixel_values = None;
- // pixel_values_video = None;
- // continue;
- // }
- // error_tokens.clear();
-
- // // 处理特殊标记和工具调用
- // match decoded_token.as_str() {
- // "" => {
- // // 开始工具调用
- // tool_call_id = Some(uuid::Uuid::new_v4().to_string());
- // seqlen_offset += seq_len;
- // seq_len = 1;
- // input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
- // cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
- // pixel_values = None;
- // pixel_values_video = None;
- // continue;
- // }
- // "" => {
- // // 结束工具调用
- // let chunk = build_completion_chunk_response(
- // decoded_token,
- // &self.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_token);
- // seqlen_offset += seq_len;
- // seq_len = 1;
- // input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
- // cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
- // pixel_values = None;
- // pixel_values_video = None;
- // continue;
- // } else {
- // // 正常文本输出
- // let chunk = build_completion_chunk_response(
- // decoded_token,
- // &self.model_name,
- // None,
- // None
- // );
- // yield Ok(chunk);
- // }
- // }
- // }
- // if next_token == self.eos_token_id1 || next_token == self.eos_token_id2 {
- // break;
- // }
- // seqlen_offset += seq_len;
- // seq_len = 1;
- // input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
- // cache_position = Tensor::from_vec(vec![seqlen_offset as u32], 1, &self.device)?;
- // pixel_values = None;
- // pixel_values_video = None;
- // }
- // self.qwen3_vl.clear_kv_cache();
- // };
- // Ok(Box::new(Box::pin(stream)))
}
}
diff --git a/src/params/chat.rs b/src/params/chat.rs
index ea314d4..0dcca33 100644
--- a/src/params/chat.rs
+++ b/src/params/chat.rs
@@ -142,6 +142,16 @@ pub struct ChatCompletionParameters {
pub top_p: Option,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_k: Option,
+ /// Penalize repetition of token sequences, helping to prevent the model from getting stuck in loops.
+ /// Values > 1.0 penalize repetitions, values < 1.0 encourage them.
+ /// For example, a value of 1.2 means that tokens that have already appeared will be less likely to appear again.
+ #[serde(skip_serializing_if = "Option::is_none")]
+ pub repeat_penalty: Option,
+ /// The number of previous tokens to consider when applying repeat penalty.
+ /// Only the last N tokens in the conversation history will be used to calculate the penalty.
+ /// A larger value considers more history for penalizing repetitions, while 0 applies no penalty.
+ #[serde(skip_serializing_if = "Option::is_none")]
+ pub repeat_last_n: Option,
/// A list of tools the model may call. Currently, only functions are supported as a tool.
/// Use this to provide a list of functions the model may generate JSON inputs for.
#[serde(skip_serializing_if = "Option::is_none")]
diff --git a/tests/test_glm_ocr.rs b/tests/test_glm_ocr.rs
index c425ba6..78ca103 100644
--- a/tests/test_glm_ocr.rs
+++ b/tests/test_glm_ocr.rs
@@ -28,7 +28,9 @@ fn glm_ocr_generate() -> Result<()> {
}
]
}
- ]
+ ],
+ "repeat_penalty": 1.2,
+ "repeat_last_n": 64
}
"#;
let save_dir =