如何基于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。
实现步骤:
- 安装依赖:
pip install transformers torch
- 批量分类代码:
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
相关产品推荐
相关产品推荐

