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

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}")

问题原因与修正方案

核心问题点

  1. 同步生成器阻塞:服务端的generate_text是同步生成器,FastAPI处理同步生成器时会等待全部内容生成完毕才返回,而非实时推送。
  2. SSE格式不规范:设置了media_type="text/event-stream"但未遵循SSE格式要求(每段内容需以data: 开头,结尾加\n\n),导致客户端缓冲直到接收完所有内容。
  3. 客户端参数错误:客户端传入的data为空,会触发服务端参数校验错误,实际测试时无法正常请求。
  4. 语法错误:服务端代码中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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.20 16:42:44