如何实现Hugging Face文本生成的用户驱动安全取消机制(无需重启服务器)
实现Hugging Face文本生成的用户驱动取消机制
下面是几种可行的安全实现方案,能在用户取消请求时中断generate()调用,同时避免服务器崩溃或GPU内存泄漏:
方案1:自定义回调+线程事件标志(跨平台通用)
利用generate()的callback参数,在每个token生成后检查取消标志,一旦触发就抛出异常终止生成。同时将生成任务放到子线程,避免阻塞主线程处理用户取消请求。
代码示例:
from transformers import pipeline, GenerationCallback import threading import torch # 定义取消事件 cancel_event = threading.Event() class CancelCallback(GenerationCallback): def on_token_generated(self, model, token_ids, **kwargs): # 每次生成token时检查取消标志 if cancel_event.is_set(): # 抛出异常终止生成,transformers会自动清理部分资源 raise RuntimeError("User requested cancellation") def generate_task(prompt, max_new_tokens): global output try: generator = pipeline('text-generation', model="TheBloke/Mistral-7B-Instruct-v0.1-GGUF", device=0) output = generator( prompt, max_new_tokens=max_new_tokens, callback=CancelCallback() ) print("生成完成:", output) except RuntimeError as e: if str(e) == "User requested cancellation": print("用户已取消生成") # 手动清理GPU内存 if torch.cuda.is_available(): torch.cuda.empty_cache() else: raise # 启动生成线程 prompt = "Tell me a very detailed history of the Roman empire..." thread = threading.Thread(target=generate_task, args=(prompt, 1000)) thread.start() # 模拟用户2秒后取消请求 import time time.sleep(2) cancel_event.set() thread.join()
关键注意事项:
- 回调函数会在每个token生成后触发,确保取消延迟不会超过单个token的生成时间
- 捕获自定义异常后,必须调用
torch.cuda.empty_cache()清理未释放的GPU张量 - 若使用
pipeline,每次生成后最好重新实例化或确保模型状态未被破坏(部分模型在异常终止后可能残留中间状态,可通过重新加载模型避免)
方案2:异步生成+Asyncio Task取消(适合Web服务场景)
如果你的服务基于异步框架(如FastAPI、Starlette),可以使用transformers的异步模型支持,结合asyncio的Task取消机制实现非阻塞的生成与取消。
代码示例:
from transformers import AutoTokenizer, AutoModelForCausalLM import asyncio import torch async def async_generate(prompt, max_new_tokens, cancel_event): tokenizer = AutoTokenizer.from_pretrained("TheBloke/Mistral-7B-Instruct-v0.1-GGUF") model = AutoModelForCausalLM.from_pretrained("TheBloke/Mistral-7B-Instruct-v0.1-GGUF", device_map="auto") inputs = tokenizer(prompt, return_tensors="pt").to(model.device) outputs = [] # 逐token生成,每步检查取消事件 for _ in range(max_new_tokens): if cancel_event.is_set(): print("用户已取消生成") # 清理资源 del inputs, model if torch.cuda.is_available(): torch.cuda.empty_cache() return None with torch.no_grad(): next_token_logits = model(**inputs).logits[:, -1, :] next_token = torch.argmax(next_token_logits, dim=-1) outputs.append(next_token.item()) inputs["input_ids"] = torch.cat([inputs["input_ids"], next_token.unsqueeze(0)], dim=-1) result = tokenizer.decode(outputs, skip_special_tokens=True) del inputs, model torch.cuda.empty_cache() return result async def main(): cancel_event = asyncio.Event() prompt = "Tell me a very detailed history of the Roman empire..." # 启动生成任务 generate_task = asyncio.create_task(async_generate(prompt, 1000, cancel_event)) # 模拟用户2秒后取消 await asyncio.sleep(2) cancel_event.set() result = await generate_task print("生成结果:", result) asyncio.run(main())
关键注意事项:
- 异步生成需要手动实现逐token生成逻辑,无法直接使用
pipeline的异步接口(部分版本的transformers异步pipeline可能不支持中途取消) - 必须在取消后手动删除模型和张量引用,并调用
torch.cuda.empty_cache()释放GPU内存 - 适合Web服务场景,可直接与FastAPI的
BackgroundTasks或async def路由结合处理用户取消请求
方案3:信号量中断(仅限Linux/macOS)
在Linux或macOS环境下,可通过发送SIGINT或自定义信号中断生成线程,但需注意信号处理的线程安全性,以及GPU资源的清理。
代码示例:
from transformers import pipeline import signal import threading import torch generator = pipeline('text-generation', model="TheBloke/Mistral-7B-Instruct-v0.1-GGUF", device=0) output = None cancel_flag = False def signal_handler(signum, frame): global cancel_flag cancel_flag = True print("用户已取消生成") # 注册信号处理函数(比如用SIGUSR1作为自定义取消信号) signal.signal(signal.SIGUSR1, signal_handler) def generate_task(prompt, max_new_tokens): global output try: # 重写generate的内部循环,加入取消检查 original_generate = generator.model.generate def interrupted_generate(*args, **kwargs): for step in range(kwargs.get("max_new_tokens", 100)): if cancel_flag: raise RuntimeError("Cancelled by user") # 调用原generate生成单个token kwargs["max_new_tokens"] = 1 result = original_generate(*args, **kwargs) args = (result,) return result generator.model.generate = interrupted_generate output = generator(prompt, max_new_tokens=max_new_tokens) print("生成完成:", output) except RuntimeError as e: if str(e) == "Cancelled by user": # 清理资源 torch.cuda.empty_cache() else: raise finally: # 恢复原generate方法 generator.model.generate = original_generate # 启动生成线程 prompt = "Tell me a very detailed history of the Roman empire..." thread = threading.Thread(target=generate_task, args=(prompt, 1000)) thread.start() # 模拟用户2秒后发送取消信号(实际场景中由UI触发服务端发送信号) import os import time time.sleep(2) os.kill(os.getpid(), signal.SIGUSR1) thread.join()
关键注意事项:
- 该方案仅适用于Linux/macOS,Windows不支持POSIX信号
- 信号处理函数必须是线程安全的,避免在信号处理中执行复杂操作
- 重写
generate方法时需确保恢复原方法,避免影响后续生成任务
通用安全准则
无论使用哪种方案,都必须遵守以下准则避免内存泄漏或崩溃:
- 及时清理GPU内存:每次取消或生成完成后,调用
torch.cuda.empty_cache(),并手动删除模型、张量等大对象引用 - 避免共享模型实例:若服务处理多个用户请求,最好为每个请求创建独立的模型实例(或使用模型池),避免一个请求的取消影响其他请求
- 捕获并处理取消异常:确保取消触发的异常被正确捕获,避免未处理的异常导致进程崩溃
内容的提问来源于stack exchange,提问作者Swati
相关产品推荐
相关产品推荐

