如何在Streamlit中实现LLM生成内容的网页流式输出?
实现Streamlit实时流式输出的解决方案
核心思路
要在Streamlit中实现类似ChatGPT的实时流式输出,需要自定义适配Streamlit的回调处理器,替代默认的命令行回调,同时让LLM开启流式模式,在生成每个token时实时更新页面内容。
步骤1:自定义Streamlit流式回调处理器
创建继承自LangChain StreamingStdOutCallbackHandler的类,重写on_llm_new_token方法,将每个生成的token实时推送到Streamlit的动态占位符中。
步骤2:修改LLM初始化配置
初始化LlamaCpp时开启streaming=True,并传入自定义的回调处理器,让LLM在生成token时触发页面更新。
步骤3:重构响应处理逻辑
使用st.empty()创建可动态更新的占位符,逐步拼接生成的token,同时将完整响应存入会话历史。
修改后的完整代码
from langchain_community.vectorstores import FAISS from langchain_community.embeddings import HuggingFaceEmbeddings from langchain import PromptTemplate from langchain_community.llms import LlamaCpp from langchain.chains import RetrievalQA from langchain.callbacks.base import BaseCallbackHandler import streamlit as st from HtmlTemplates import bot_template, user_template, css import torch # 自定义Streamlit流式回调处理器 class StreamlitStreamingCallback(BaseCallbackHandler): def __init__(self, placeholder): self.placeholder = placeholder self.response_text = "" def on_llm_new_token(self, token: str, **kwargs) -> None: self.response_text += token # 用bot模板包裹实时更新的内容 self.placeholder.markdown(bot_template.replace("{{MSG}}", self.response_text), unsafe_allow_html=True) def set_prompt(): custom_prompt_template = """[INST] <<SYS>> You are a trained to guide people about Indian Law. You will answer user's query with your knowledge and use context provided. Do not say thank you and tell you are an AI Assistant and be open about everything. Always complete the sentence you are generating <</SYS>> Use the following pieces of context to answer the users question. Context : {context} Question : {question} Answer : [/INST] """ prompt = PromptTemplate(template=custom_prompt_template, input_variables=["context", "question"]) return prompt def retrieval_qa_chain(llm, prompt, db): qa_chain = RetrievalQA.from_chain_type( llm=llm, chain_type='stuff', retriever=db.as_retriever(search_kwargs={'k': 6}), chain_type_kwargs={'prompt': prompt} ) return qa_chain def qa_pipeline(streaming_callback=None): device = torch.device("cuda" if torch.cuda.is_available() else "cpu") embeddings = HuggingFaceEmbeddings(model_name='multi-qa-mpnet-base-dot-v1', model_kwargs={'device': device}) db = FAISS.load_local("vectorstore", embeddings, allow_dangerous_deserialization=True) # 初始化LlamaCpp时开启流式,并传入自定义回调 llm = LlamaCpp( model_path=path, temperature=temperature, n_ctx=2048, n_batch=128, n_gpu_layers=-1, max_tokens=max_tokens, verbose=False, streaming=True, # 开启流式模式 callbacks=[streaming_callback] if streaming_callback else [] # 传入自定义回调 ) qa_prompt = set_prompt() chain = retrieval_qa_chain(llm, qa_prompt, db) return chain def handle_user_input(user_question): # 创建空占位符用于实时更新响应 response_placeholder = st.empty() # 初始化自定义回调,绑定占位符 stream_callback = StreamlitStreamingCallback(response_placeholder) # 重新初始化带回调的chain chain = qa_pipeline(stream_callback) with st.spinner("Generating response ..."): # 执行查询,回调会实时更新页面 response = chain(user_question) full_response = response['result'] # 将完整响应存入会话历史 st.session_state.chat_history.append({"User": user_question, "Bot": full_response}) st.set_page_config(page_title="Your personal Law ChatBot", page_icon=":bot:") st.write(css, unsafe_allow_html=True) global chain, path, temperature, max_tokens with st.sidebar: model = st.selectbox("Select Model :", ("Llama2 7b (Faster)", "Llama2 13b (Can answer complex queries)")) if model == 'Llama2 13b (Can answer complex queries)': path = "Models/llama-2-13b-chat.Q4_K_M.gguf" elif model == 'Llama2 7b (Faster)': path = "Models/llama-2-7b-chat.Q4_K_M.gguf" temperature = st.slider( label="Temperature", min_value=0.1, max_value=1.0, value=0.3, step=0.05 ) max_tokens = st.slider( label="Max Tokens", min_value=256, max_value=4096, value=1024, step=64 ) if "chat_history" not in st.session_state: st.session_state.chat_history = [] st.header("Your personal Law ChatBot :books:") user_question = st.chat_input("Ask a question :") if user_question: # 显示用户提问 st.write(user_template.replace("{{MSG}}", user_question), unsafe_allow_html=True) # 处理流式响应 handle_user_input(user_question) # 显示历史会话 for chat in st.session_state.chat_history: st.write(user_template.replace("{{MSG}}", chat["User"]), unsafe_allow_html=True) st.write(bot_template.replace("{{MSG}}", chat["Bot"]), unsafe_allow_html=True)
关键修改说明
- 自定义回调处理器:
StreamlitStreamingCallback类负责接收每个新生成的token,实时更新Streamlit占位符内容,实现动态刷新效果。 - LLM流式模式:在
LlamaCpp初始化时设置streaming=True,并传入自定义回调,让LLM生成token时触发页面更新。 - 动态占位符:使用
st.empty()创建可更新区域,替代原一次性输出的st.write,实现内容逐步显示。
内容的提问来源于stack exchange,提问作者Ashish Sawant
相关产品推荐
相关产品推荐

