如何中断Transformers LLM的generate()函数?RAG聊天Bot场景
解决方案
方案1:利用LangChain回调函数实现协作式中断
LangChain的生成流程支持自定义回调,你可以在每生成一个token的节点检查Streamlit的中断信号,一旦检测到用户触发中断,就抛出异常终止生成。
实现步骤
- 在Streamlit会话状态中维护一个中断标记(如
st.session_state['interrupt']),通过按钮触发将其设为True。 - 自定义LangChain的
CallbackHandler,在on_llm_new_token方法中检查中断标记,若触发则抛出KeyboardInterrupt终止生成。 - 调用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
相关产品推荐
相关产品推荐

