Mistral-7B流式输出仅服务端生效,客户端获完整输出问题排查
服务端流式返回失效,客户端仅收到完整结果排查
我搭了一套服务端-客户端远程流水线,预期服务端加载Mistral-7B模型后流式返回生成结果。现在服务端内部能正常流式输出,但客户端只能拿到完整的生成结果,问题出在哪?
服务端代码
from fastapi import FastAPI, Request, HTTPException from fastapi.responses import StreamingResponse from transformers import AutoModelForCausalLM, AutoTokenizer, TextStreamer import torch from pydantic import BaseModel import sys import os from queue import Queue from threading import Thread from inference import set_model, process_audio_streaming app = FastAPI() # Load the tokenizer and model model, tokenizer, args = set_model() # replace with your script class PredictionRequest(BaseModel): raw_feedback: str class PredictionResponse(BaseModel): prediction: str class CustomTextStreamer(TextStreamer): def __init__(self, tokenizer): super().__init__(tokenizer) self.queue = Queue() def on_text(self, text: str, **kwargs): self.queue.put(text) def get_generated_text(self): while True: text = self.queue.get() if text is None: break yield text def generate_text(prompt: str, max_new_tokens: int = 256): device = "cuda" if torch.cuda.is_available() else "cpu" model.to(device) inputs = tokenizer(prompt, return_tensors="pt").to(device) streamer = CustomTextStreamer(tokenizer) def generate(): model.generate(inputs['input_ids'], streamer=streamer, max_new_tokens=max_new_tokens) streamer.queue.put(None) # Signal the end of generation generation_thread = Thread(target=generate) generation_thread.start() return streamer.get_generated_text() @app.post("/generate-text") async def generate_text_endpoint(request: Request): body = await request.json() raw_feedback = body.get("raw_feedback") if raw_feedback is None is None: raise HTTPException(status_code=400, detail="raw_feedback is required") # Process the raw feedback and accuracy # raw_feedback = process_audio_streaming(raw_feedback, args) # Define the prompt using the processed feedback and accuracy prompt = ( "tet prompt" ) # Stream the text as it's generated return StreamingResponse(generate_text(prompt), media_type="text/event-stream") #text/plain if __name__ == "__main__": import uvicorn uvicorn.run(app, host="0.0.0.0", port=8000)
客户端代码
import requests # URL of the FastAPI endpoint url = 'http://localhost:8000/generate-text' # Data to be sent to the endpoint data = {} try: with requests.post(url, json=data, stream=True) as r: r.raise_for_status() # Ensure we catch HTTP errors for chunk in r.iter_content(chunk_size=1024): if chunk: print(chunk.decode('utf-8'), end='', flush=True) except requests.RequestException as e: print(f"An error occurred: {e}")
问题原因与修正方案
核心问题点
- 同步生成器阻塞:服务端的
generate_text是同步生成器,FastAPI处理同步生成器时会等待全部内容生成完毕才返回,而非实时推送。 - SSE格式不规范:设置了
media_type="text/event-stream"但未遵循SSE格式要求(每段内容需以data:开头,结尾加\n\n),导致客户端缓冲直到接收完所有内容。 - 客户端参数错误:客户端传入的
data为空,会触发服务端参数校验错误,实际测试时无法正常请求。 - 语法错误:服务端代码中
if raw_feedback is None is None:是无效语法,应改为if raw_feedback is None:。
修正后的服务端代码
将同步生成器改为异步适配,同时规范SSE输出格式:
from fastapi import FastAPI, Request, HTTPException from fastapi.responses import StreamingResponse from transformers import AutoModelForCausalLM, AutoTokenizer, TextStreamer import torch from pydantic import BaseModel import sys import os from queue import Queue from threading import Thread import asyncio from inference import set_model, process_audio_streaming app = FastAPI() # Load the tokenizer and model model, tokenizer, args = set_model() # replace with your script class PredictionRequest(BaseModel): raw_feedback: str class PredictionResponse(BaseModel): prediction: str class CustomTextStreamer(TextStreamer): def __init__(self, tokenizer): super().__init__(tokenizer) self.queue = Queue() def on_text(self, text: str, **kwargs): self.queue.put(text) async def async_generate_text(prompt: str, max_new_tokens: int = 256): device = "cuda" if torch.cuda.is_available() else "cpu" model.to(device) inputs = tokenizer(prompt, return_tensors="pt").to(device) streamer = CustomTextStreamer(tokenizer) def generate(): model.generate(inputs['input_ids'], streamer=streamer, max_new_tokens=max_new_tokens) streamer.queue.put(None) # Signal the end of generation # 启动线程执行同步生成逻辑 generation_thread = Thread(target=generate) generation_thread.start() # 异步读取队列,按SSE格式输出 while True: if not streamer.queue.empty(): text = streamer.queue.get() if text is None: break yield f"data: {text}\n\n" await asyncio.sleep(0.01) # 避免空轮询占用资源 @app.post("/generate-text") async def generate_text_endpoint(request: Request): body = await request.json() raw_feedback = body.get("raw_feedback") if raw_feedback is None: raise HTTPException(status_code=400, detail="raw_feedback is required") # Process the raw feedback and accuracy # raw_feedback = process_audio_streaming(raw_feedback, args) # Define the prompt using the processed feedback and accuracy prompt = "test prompt" # 修正原代码笔误 # Stream the text as it's generated return StreamingResponse(async_generate_text(prompt), media_type="text/event-stream") if __name__ == "__main__": import uvicorn # 启动时关闭服务端缓冲,确保流式输出 uvicorn.run(app, host="0.0.0.0", port=8000, timeout_keep_alive=60)
修正后的客户端代码
适配SSE格式解析,传入合法参数:
import requests # URL of the FastAPI endpoint url = 'http://localhost:8000/generate-text' # Data to be sent to the endpoint data = {"raw_feedback": "你的测试输入内容"} try: with requests.post(url, json=data, stream=True, headers={"Accept": "text/event-stream"}) as r: r.raise_for_status() # 逐行解析SSE内容 for line in r.iter_lines(decode_unicode=True): if line.startswith('data: '): content = line[6:].strip() print(content, end='', flush=True) except requests.RequestException as e: print(f"An error occurred: {e}")
内容的提问来源于stack exchange,提问作者Phys
相关产品推荐
相关产品推荐

