无需TGI实现HuggingFace流式推理:GPU容器问题与前端输出需求
实现大模型流式输出到前端的替代方案
方案1:Hugging Face Transformers + FastAPI 自定义流式接口
直接用Transformers库结合FastAPI的流式响应能力实现,完全可控,无需依赖TGI镜像。
- 安装依赖:
pip install transformers torch fastapi uvicorn
- 核心代码示例:
from fastapi import FastAPI, Request from fastapi.responses import StreamingResponse from transformers import AutoTokenizer, AutoModelForCausalLM, TextIteratorStreamer import torch from threading import Thread app = FastAPI() # 替换为你的模型名称 tokenizer = AutoTokenizer.from_pretrained("your-model-name") model = AutoModelForCausalLM.from_pretrained("your-model-name", torch_dtype=torch.float16).to("cuda") @app.post("/stream-chat") async def stream_chat(request: Request): data = await request.json() prompt = data["prompt"] # 初始化流式迭代器,跳过输入提示内容 streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, timeout=30.0) # 构建模型输入 inputs = tokenizer(prompt, return_tensors="pt").to("cuda") # 用线程运行模型生成,避免阻塞主线程 generation_kwargs = dict(inputs, streamer=streamer, max_new_tokens=512) thread = Thread(target=model.generate, kwargs=generation_kwargs) thread.start() # 定义生成器,逐块返回文本给前端 def generate(): for new_text in streamer: yield f"data: {new_text}\n\n" return StreamingResponse(generate(), media_type="text/event-stream")
- 前端通过
EventSource或Fetch API的ReadableStream接收SSE(Server-Sent Events),即可实现打字机式的流式输出效果。
方案2:使用vLLM部署流式接口
vLLM是高性能大模型推理引擎,原生支持流式输出,GPU利用率优于原生Transformers,部署流程简单。
- 安装vLLM:
pip install vllm
- 启动流式API服务:
python -m vllm.entrypoints.api_server --model your-model-name --gpu-memory-utilization 0.9 --enable-streaming
- 前端调用
/v1/completions或/v1/chat/completions接口时,设置stream: true,即可接收与OpenAI API格式兼容的流式返回结果,前端代码可直接复用ChatGPT相关逻辑。
方案3:自定义TextStreamer重定向输出
如果不想更换框架,可以重写TextStreamer的输出逻辑,将文本发送至前端而非标准输出。
示例代码:
from transformers import TextStreamer import asyncio class WebStreamer(TextStreamer): def __init__(self, tokenizer, queue): super().__init__(tokenizer) self.queue = queue def on_finalized_text(self, text: str, stream_end: bool = False): # 将生成的文本放入异步队列,供接口读取后返回给前端 asyncio.run_coroutine_threadsafe(self.queue.put(text), asyncio.get_event_loop()) if stream_end: asyncio.run_coroutine_threadsafe(self.queue.put(None), asyncio.get_event_loop()) # 在FastAPI接口中通过异步队列读取WebStreamer的输出,再返回给前端
内容的提问来源于stack exchange,提问作者Muhammad Fhadli
相关产品推荐
相关产品推荐

