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

如何实现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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 07:20:03