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

如何基于HuggingFace的BART模型加速零样本新闻标题分类?

大幅提升新闻标题分类速度的解决方案

1. 批量调用Hugging Face API(最快实现,改动最小)

远程API的核心瓶颈是多次网络请求的开销,将300条标题一次性批量传入,把300次请求压缩为1次,能直接将总耗时降到API处理批量数据的时间(通常几秒内完成),完全满足你的需求。

修改后的代码示例:

import requests

def batch_categorise(categories, items):
    API_URL = "https://api-inference.huggingface.co/models/facebook/bart-large-mnli"
    headers = {"Authorization": "Bearer <token>"}
    payload = {
        "inputs": items,  # 传入标题列表,API支持批量处理多条目
        "parameters": {"candidate_labels": categories},
    }
    response = requests.post(API_URL, headers=headers, json=payload)
    if response.status_code == 200:
        results = response.json()
        return [
            {"category": res['labels'][0], "confidence": res['scores'][0]}
            for res in results
        ]
    else:
        return {"error": f"API request failed with status {response.status_code}"}

# 使用方式
titles = ["标题1", "标题2", ..., "标题300"]
categories = ["类别A", "类别B", ...]
results = batch_categorise(categories, titles)

2. 本地部署模型(性能最优,长期方案)

远程API的网络延迟是不可避免的瓶颈,本地部署模型能彻底消除这个问题,结合GPU加速后处理速度会远超远程API。

实现步骤:

  1. 安装依赖:
pip install transformers torch
  1. 批量分类代码:
from transformers import pipeline

# 加载模型(首次运行自动下载,后续直接读取本地缓存)
classifier = pipeline("zero-shot-classification", model="facebook/bart-large-mnli", device=0)  # device=0启用GPU,无GPU则删除该参数

def batch_local_categorise(categories, items):
    results = classifier(items, candidate_labels=categories)
    return [
        {"category": res['labels'][0], "confidence": res['scores'][0]}
        for res in results
    ]

# 使用方式
titles = ["标题1", "标题2", ..., "标题300"]
categories = ["类别A", "类别B", ...]
results = batch_local_categorise(categories, titles)
  • 有GPU的情况下,300条标题处理时间在1-3秒内;即使仅用CPU,也能在5秒左右完成。
  • 若想进一步提速,可替换为轻量化模型(如distilbart-mnli-12-3,速度提升3倍左右),或用bitsandbytes库对模型做4bit量化,减少内存占用的同时加快推理速度。

3. 并发异步请求(保留远程API的折中方案)

如果无法本地部署模型,用异步并发请求替代同步循环,能大幅缩短总耗时(效果弱于批量调用)。

代码示例:

import aiohttp
import asyncio

async def async_categorise(session, categories, item):
    API_URL = "https://api-inference.huggingface.co/models/facebook/bart-large-mnli"
    headers = {"Authorization": "Bearer <token>"}
    payload = {
        "inputs": item,
        "parameters": {"candidate_labels": categories},
    }
    async with session.post(API_URL, headers=headers, json=payload) as response:
        if response.status == 200:
            result = await response.json()
            return {"category": result[0]['labels'][0], "confidence": result[0]['scores'][0]}
        else:
            return {"error": f"API request failed with status {response.status}"}

async def batch_async_categorise(categories, items):
    # 用信号量控制并发数,避免触发API限流
    semaphore = asyncio.Semaphore(20)
    async def bounded_categorise(item):
        async with semaphore:
            return await async_categorise(session, categories, item)
    
    async with aiohttp.ClientSession() as session:
        tasks = [bounded_categorise(item) for item in items]
        return await asyncio.gather(*tasks)

# 使用方式
titles = ["标题1", "标题2", ..., "标题300"]
categories = ["类别A", "类别B", ...]
results = asyncio.run(batch_async_categorise(categories, titles))

额外优化建议

  • 精简候选类别数量:去掉重复或冗余的类别,能降低模型计算量,提升处理速度。
  • 服务器区域优化:若Flask部署在云服务器,选择靠近Hugging Face API节点的区域(如美国东部、欧洲),减少网络延迟。

内容的提问来源于stack exchange,提问作者karpathy

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 00:06:04