DeepSeek-R1-Distill-Qwen-1.5B实战教程:添加自定义stop_token应对无限生成边界问题

1. 项目简介

DeepSeek-R1-Distill-Qwen-1.5B是一个超轻量级的智能对话模型,它在魔塔平台上获得了极高的下载量。这个模型巧妙结合了DeepSeek优秀的逻辑推理能力和Qwen成熟的模型架构,经过蒸馏优化后,在保持强大能力的同时大幅降低了计算资源需求。

这个1.5B参数的模型特别适合在低显存GPU或轻量计算环境中运行,让你不需要昂贵的硬件就能享受到高质量的AI对话体验。项目基于Streamlit构建了直观的可视化界面,支持多轮对话、思维链推理,并且所有处理都在本地完成,完全保障你的数据隐私。

2. 无限生成问题的由来

在实际使用中,你可能会遇到一个常见问题:模型有时候会停不下来,一直生成内容直到达到最大长度限制。这种情况通常发生在模型无法准确识别何时应该结束回答的时候。

2.1 为什么会出现无限生成

模型生成文本的原理是基于概率预测下一个token,当模型不确定何时结束时,它会继续生成,导致回答变得冗长甚至偏离主题。特别是在处理复杂推理任务时,模型可能会陷入循环思考或者不断补充细节。

2.2 传统停止条件的局限性

大多数模型使用简单的停止条件,比如遇到句号、问号等标点符号,或者达到最大生成长度。但这些方法往往不够智能,无法准确判断在什么语境下应该停止生成。

3. 自定义stop_token解决方案

为了解决无限生成的问题,我们可以为模型添加自定义的停止标记(stop_token),让模型在生成特定内容时自动停止。

3.1 理解stop_token机制

stop_token就像是给模型设置的红绿灯,当模型生成这些特定的标记时,就会自动停止生成。这比单纯依赖标点符号或长度限制要智能得多。

# 基本的stop_token使用示例
stop_tokens = ["思考完毕", "回答结束", "[END]"]

def generate_with_stop_tokens(prompt, stop_tokens):
    # 这里是生成的逻辑
    # 当模型输出中包含stop_tokens中的任何一个时,就停止生成
    pass

3.2 为DeepSeek-R1模型添加stop_token

针对DeepSeek-R1模型的特性,我们可以设置一些智能的停止标记:

def setup_custom_stop_tokens():
    """为DeepSeek-R1模型设置自定义停止标记"""
    custom_stop_tokens = [
        "\n思考完毕",
        "\n回答结束", 
        "\n综上",
        "\n总结来说",
        "\n希望以上回答",
        "[END]",
        "<|endoftext|>"
    ]
    
    # 将这些停止标记添加到生成配置中
    return custom_stop_tokens

4. 实战配置步骤

现在让我们一步步实现为DeepSeek-R1模型添加自定义stop_token的功能。

4.1 修改模型生成参数

首先,我们需要修改模型的生成配置,加入我们的停止标记:

from transformers import AutoTokenizer, AutoModelForCausalLM
import torch

# 加载模型和分词器
model_path = "/root/ds_1.5b"
tokenizer = AutoTokenizer.from_pretrained(model_path)
model = AutoModelForCausalLM.from_pretrained(
    model_path,
    torch_dtype=torch.float16,
    device_map="auto"
)

# 设置自定义停止标记
def get_generation_config():
    """获取包含自定义stop_token的生成配置"""
    stop_tokens = ["思考完毕", "回答结束", "\n综上", "\n总结"]
    
    # 将停止标记转换为token ID
    stop_token_ids = []
    for token in stop_tokens:
        stop_token_ids.extend(tokenizer.encode(token, add_special_tokens=False))
    
    generation_config = {
        "max_new_tokens": 2048,
        "temperature": 0.6,
        "top_p": 0.95,
        "do_sample": True,
        "stop_token_ids": stop_token_ids,
        "pad_token_id": tokenizer.eos_token_id
    }
    
    return generation_config

4.2 实现智能停止检测

单纯的停止标记可能还不够,我们需要更智能的停止检测机制:

def smart_stopping_criteria(generated_text, stop_tokens):
    """智能停止条件检测"""
    # 检查是否包含任何停止标记
    for stop_token in stop_tokens:
        if stop_token in generated_text:
            return True
    
    # 检查是否已经给出了完整的回答
    if is_complete_answer(generated_text):
        return True
        
    return False

def is_complete_answer(text):
    """判断是否已经形成完整回答"""
    # 检查是否有明显的结束信号
    end_indicators = ["。", "!", "?", "\n\n", "以上就是"]
    for indicator in end_indicators:
        if text.endswith(indicator):
            return True
    
    # 检查是否有总结性语句
    summary_phrases = ["总之", "综上所述", "总的来说", "因此"]
    for phrase in summary_phrases:
        if phrase in text[-20:]:  # 检查最后20个字符
            return True
            
    return False

5. 集成到Streamlit应用

现在我们将这些功能集成到Streamlit聊天应用中。

5.1 修改生成函数

import streamlit as st
from transformers import StoppingCriteria, StoppingCriteriaList

class CustomStoppingCriteria(StoppingCriteria):
    """自定义停止条件"""
    def __init__(self, stop_token_ids):
        self.stop_token_ids = stop_token_ids
    
    def __call__(self, input_ids, scores, **kwargs):
        # 检查最近生成的token是否在停止标记中
        last_token = input_ids[0][-1].item()
        return last_token in self.stop_token_ids

def generate_response_with_stop(prompt, conversation_history):
    """使用自定义停止标记生成回答"""
    # 准备停止条件
    stop_tokens = ["思考完毕", "回答结束", "\n综上"]
    stop_token_ids = [tokenizer.encode(token, add_special_tokens=False)[0] for token in stop_tokens]
    stopping_criteria = StoppingCriteriaList([CustomStoppingCriteria(stop_token_ids)])
    
    # 准备输入
    chat_template = tokenizer.apply_chat_template(
        conversation_history,
        tokenize=False,
        add_generation_prompt=True
    )
    
    inputs = tokenizer(chat_template, return_tensors="pt").to(model.device)
    
    # 生成回答
    with torch.no_grad():
        outputs = model.generate(
            **inputs,
            max_new_tokens=2048,
            temperature=0.6,
            top_p=0.95,
            do_sample=True,
            stopping_criteria=stopping_criteria,
            pad_token_id=tokenizer.eos_token_id
        )
    
    # 解码并处理输出
    response = tokenizer.decode(outputs[0][inputs.input_ids.shape[1]:], skip_special_tokens=True)
    return process_response(response)

def process_response(response):
    """处理模型输出,格式化思考过程"""
    # 替换思考标签为更友好的格式
    response = response.replace("<|im_start|>assistant\n", "")
    response = response.replace("<|im_end|>", "")
    
    # 自动截断在停止标记之后的内容
    stop_points = ["思考完毕", "回答结束", "\n综上"]
    for stop_point in stop_points:
        if stop_point in response:
            response = response.split(stop_point)[0] + stop_point
            break
    
    return response

5.2 更新聊天界面逻辑

# 在Streamlit应用中集成新功能
if "messages" not in st.session_state:
    st.session_state.messages = []

# 显示聊天记录
for message in st.session_state.messages:
    with st.chat_message(message["role"]):
        st.markdown(message["content"])

# 处理用户输入
if prompt := st.chat_input("考考 DeepSeek R1..."):
    st.session_state.messages.append({"role": "user", "content": prompt})
    
    with st.chat_message("user"):
        st.markdown(prompt)
    
    # 生成回答
    with st.chat_message("assistant"):
        with st.spinner("思考中..."):
            response = generate_response_with_stop(prompt, st.session_state.messages)
            st.markdown(response)
    
    st.session_state.messages.append({"role": "assistant", "content": response})

6. 效果测试与优化

添加自定义stop_token后,让我们测试一下效果并进一步优化。

6.1 测试不同场景下的停止效果

def test_stopping_effectiveness():
    """测试停止标记在不同场景下的效果"""
    test_cases = [
        "请解释什么是机器学习",
        "帮我写一个Python函数计算斐波那契数列",
        "分析一下气候变化对农业的影响",
        "请用思维链的方式解决这个数学问题:2+2=?"
    ]
    
    for test_case in test_cases:
        print(f"测试问题: {test_case}")
        response = generate_response_with_stop(test_case, [])
        print(f"生成结果: {response}")
        print("-" * 50)

# 运行测试
test_stopping_effectiveness()

6.2 优化停止策略

根据测试结果,我们可以进一步优化停止策略:

def optimize_stopping_strategy():
    """基于测试结果优化停止策略"""
    # 收集常见的结束模式
    common_end_patterns = [
        # 中文结束模式
        "。", "!", "?", "……",
        # 总结性短语
        "总之", "综上所述", "总的来说", "因此",
        # 礼貌性结束
        "希望以上回答", "如果还有其他问题", "谢谢提问",
        # 思维链结束
        "思考完毕", "推理完成", "解答结束"
    ]
    
    # 根据模型输出特点调整停止标记
    optimized_stop_tokens = common_end_patterns + [
        "\n\n",  # 空行通常表示内容结束
        "[END]",
        "<|endoftext|>"
    ]
    
    return optimized_stop_tokens

7. 总结

通过为DeepSeek-R1-Distill-Qwen-1.5B模型添加自定义stop_token,我们有效解决了无限生成的问题。这个方案不仅让模型回答更加简洁精准,还提升了用户体验。

7.1 主要收获

  1. 精准控制生成长度:模型现在能够在合适的时机停止生成,避免冗长回答
  2. 提升回答质量:停止标记帮助模型产出更加结构化和完整的回答
  3. 更好的用户体验:用户不再需要手动截断过长的回答

7.2 实践建议

  • 根据你的具体使用场景调整停止标记
  • 定期测试和优化停止策略
  • 结合最大生成长度限制,形成双重保障
  • 监控模型输出,及时发现新的停止模式

7.3 进一步优化方向

未来可以考虑使用更智能的停止策略,比如基于语义理解来判断回答是否完整,或者使用机器学习方法来预测最佳停止点。这样可以让模型的停止行为更加自然和智能。


获取更多AI镜像

想探索更多AI镜像和应用场景?访问 CSDN星图镜像广场,提供丰富的预置镜像,覆盖大模型推理、图像生成、视频生成、模型微调等多个领域,支持一键部署。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐