You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何中断Transformers LLM的generate()函数?RAG聊天Bot场景

解决方案

方案1:利用LangChain回调函数实现协作式中断

LangChain的生成流程支持自定义回调,你可以在每生成一个token的节点检查Streamlit的中断信号,一旦检测到用户触发中断,就抛出异常终止生成。

实现步骤

  1. 在Streamlit会话状态中维护一个中断标记(如st.session_state['interrupt']),通过按钮触发将其设为True。
  2. 自定义LangChain的CallbackHandler,在on_llm_new_token方法中检查中断标记,若触发则抛出KeyboardInterrupt终止生成。
  3. 调用LLM时传入该回调函数,捕获异常后给出中断提示。

代码示例

import streamlit as st
from langchain.callbacks.base import BaseCallbackHandler
from langchain.llms import HuggingFacePipeline
from transformers import pipeline

# 初始化中断状态
if 'interrupt' not in st.session_state:
    st.session_state['interrupt'] = False

# 自定义中断回调
class InterruptCallback(BaseCallbackHandler):
    def on_llm_new_token(self, token: str, **kwargs) -> None:
        if st.session_state['interrupt']:
            raise KeyboardInterrupt("用户触发中断")

# 加载LLM示例
pipe = pipeline("text-generation", model="your-model-name")
llm = HuggingFacePipeline(pipeline=pipe)

# Streamlit界面
st.title("RAG聊天机器人")
user_query = st.text_input("输入你的问题:")

if st.button("发送"):
    st.session_state['interrupt'] = False
    with st.spinner("生成回答中..."):
        try:
            response = llm(user_query, callbacks=[InterruptCallback()])
            st.write(response)
        except KeyboardInterrupt:
            st.warning("生成已被用户中断")

if st.button("中断生成"):
    st.session_state['interrupt'] = True

方案2:直接用Transformers的stopping_criteria

如果直接调用Transformers的generate方法,可以自定义StoppingCriteria类,在每一步生成时检查中断信号,满足条件就停止生成。

代码示例

import streamlit as st
from transformers import AutoModelForCausalLM, AutoTokenizer, StoppingCriteria, StoppingCriteriaList
import torch

# 初始化中断状态
if 'interrupt' not in st.session_state:
    st.session_state['interrupt'] = False

# 自定义停止条件
class InterruptStoppingCriteria(StoppingCriteria):
    def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor, **kwargs) -> bool:
        return st.session_state['interrupt']

# 加载模型和tokenizer
model_name = "your-model-name"
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForCausalLM.from_pretrained(model_name)

# Streamlit界面
st.title("RAG聊天机器人")
user_query = st.text_input("输入你的问题:")

if st.button("发送"):
    st.session_state['interrupt'] = False
    with st.spinner("生成回答中..."):
        try:
            inputs = tokenizer(user_query, return_tensors="pt")
            stopping_criteria = StoppingCriteriaList([InterruptStoppingCriteria()])
            outputs = model.generate(
                **inputs,
                stopping_criteria=stopping_criteria,
                max_new_tokens=512
            )
            response = tokenizer.decode(outputs[0], skip_special_tokens=True)
            st.write(response)
        except Exception:
            st.warning("生成已被用户中断")

if st.button("中断生成"):
    st.session_state['interrupt'] = True

关键说明

  • 线程方案失效原因:Python线程受GIL限制,无法强制终止,只能通过协作式中断让生成过程主动检查信号并停止,上述两个方案均基于此思路。
  • 内存效率:两个方案均在主进程内运行,无需额外加载LLM,避免了子进程的内存浪费问题。

内容的提问来源于stack exchange,提问作者jonby

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.29 01:07:27