codelion/adaptive-classifier · 上手攻略
- 仓库:codelion/adaptive-classifier
- 链接:https://github.com/codelion/adaptive-classifier
- 分类:ML / 文本分类
- 作者:Tom
- 更新:2026-10-09
这是什么
Adaptive Classifier 是一个基于 PyTorch 和 HuggingFace Transformers 的动态文本分类库,核心理念是"分类器可以不断学习新类别而不遗忘旧知识",并且内置了博弈论对抗防御(Strategic Classification)机制,在用户试图通过修改输入文本操纵分类结果时仍能保持鲁棒性。⚠️ 对抗数据集(AI-Secure/adv_glue)上实测:普通分类器遭遇对抗输入准确率下跌 20 个百分点,Strategic Classifier 仅下跌 0 个百分点。
特点:可接入任意 HuggingFace transformer 模型(RAGTruth 基准测试综合 F1 51.54% / Arena-Hard-Auto v0.1 效率提升 27%);内置 ONNX Runtime 加速(CPU 推理 2-4 倍);动态增删类别无需全量重训。
解决什么问题
- 类别不断涌现的生产环境:客服工单分类新类别("隐私合规"、"集成支持")出现时不需要重新收集全部类别数据重训
- 需要对抗输入操纵:用户可能故意改写文本绕过分类(如恶意投诉伪装成技术问题),Strategic 模式提供博弈论防御
- 资源有限的部署场景:ONNX 加速 CPU 推理,适合无法上 GPU 的边缘/服务端环境
- 需要持续学习不遗忘:传统 fine-tune 增数据会覆盖旧类别,Adaptive Classifier 通过 EWC 保护+原型记忆避免灾难性遗忘
典型应用:客服工单分类、LLM 输出真实性检测(幻觉检测)、模型路由(Query Classification)、RAG 检索质量评估(RAGTruth)、多语言情感分析。
快速安装
pip install adaptive-classifier
安装包自带 ONNX Runtime,CPU 推理开箱即用。如需开发模式并包含测试依赖:
git clone https://github.com/codelion/adaptive-classifier.git
cd adaptive-classifier
pip install -e ".[test]"
⚠️ 最低依赖:Python 3.9+,PyTorch 2.0+,transformers。FAISS 用于原型最近邻搜索(pip install faiss-cpu 或 faiss-gpu)。
核心用法
基本分类
from adaptive_classifier import AdaptiveClassifier
# 初始化,指定任意 HuggingFace transformer
classifier = AdaptiveClassifier("sentence-transformers/all-MiniLM-L6-v2")
# 动态添加训练示例
texts = ["Great product, highly recommend!", "Completely broken, refund please", "API returning 500 errors"]
labels = ["positive", "negative", "technical"]
classifier.add_examples(texts, labels)
# 即时分类
predictions = classifier.predict("This is amazing!")
# Returns: [('positive', 0.87), ('negative', 0.08), ('technical', 0.05)]
池化策略选择
⚠️ 重要:0.2.0 之前 CLS token 始终被使用,但这对大多数 encoder 是错误的(比如 all-MiniLM-L6-v2 是 mean-pooling 模型,CLS 在语义相似度测试中打出了 0.53 的错误相似度分)。从 0.2.0 起默认 auto,自动读取模型的 sentence-transformers 配置。
# 使用模型原配置(默认,推荐)
classifier = AdaptiveClassifier("sentence-transformers/all-MiniLM-L6-v2")
# 手动指定
classifier = AdaptiveClassifier("bert-base-uncased", config={"pooling": "mean"}) # 或 "cls"
⚠️ 在此版本之前保存的分类器重新加载后会保留 CLS 池化,如需新行为需显式传 pooling 参数或重新训练。
原型权重 vs 神经头权重
预测信号由两部分混合:原型分数(文本与各类别原型向量的余弦相似度)和神经头分数(小型可训练网络)。默认比例原型 0.7 / 神经头 0.3,可调:
classifier = AdaptiveClassifier(
"sentence-transformers/all-MiniLM-L6-v2",
config={"prototype_weight": 0.3, "neural_weight": 0.7},
)
类别少示例时的处理
⚠️ 新类别示例数少于 new_class_example_threshold(默认 10)时,头权重从 0 线性提升至 neural_weight,原型权重相应递减——确保少量样本时以原型(相似度判断)为主、神经头(过拟合风险)被压制。0.3.0 之前使用固定 0.3/0.7 比例,少量数据下神经头容易过拟合。
原型锐度调节
原型分数 = softmax(-distance² / prototype_temperature),默认 temperature = 0.25。调低则最近类别更突出,调至 None 恢复 0.3.0 之前的旧评分方式(最近类与其他的区分度很低):
classifier = AdaptiveClassifier("bert-base-uncased", config={"prototype_temperature": 0.1})
战略分类模式(对抗防御)
from adaptive_classifier import AdaptiveClassifier, PredictionMode
# 普通分类
predictions = classifier.predict("Great product!")
# 战略分类(对抗输入防御)
predictions_strategic = classifier.predict("Great product!", mode=PredictionMode.STRATEGIC)
# 鲁棒分类(分布外稳健)
predictions_robust = classifier.predict("Great product!", mode=PredictionMode.ROBUST)
在 AI-Secure/adv_glue 对抗 SST-2 数据集上:普通分类器准确率 60%,Strategic 分类器维持 82.22%(提升 +22.22pp),Robust 分类器下跌 0pp。
推送/拉取 HuggingFace 模型
# 推送到 Hub
classifier.push_to_hub("your-username/your-classifier")
# 从 Hub 拉取
classifier = AdaptiveClassifier.from_pretrained("your-username/your-classifier")
与 LangChain 集成
from langchain_core.language_models import BaseLLM
from adaptive_classifier import AdaptiveClassifier
# LangChain 的 LLM 输出后接分类器做幻觉检测
classifier = AdaptiveClassifier("bert-base-uncased")
classifier.add_examples(["factually correct response"], labels=["faithful"])
classifier.add_examples(["made up specific numbers and dates"], labels=["hallucinated"])
# 用分类器评估 LLM 输出
llm_output = "The company revenue was 2.3 billion in 2024."
result = classifier.predict(llm_output)
⚠️ 以上为示意性代码,实测幻觉检测建议参考 RAGTruth 基准测试(F1 51.54%)评估是否满足需求。
典型适用场景
| 场景 | 推荐模式 | 说明 |
|---|---|---|
| 客服工单动态分类 | 普通模式 | 持续增删类别 |
| LLM 输出幻觉检测 | Strategic 或 Robust | 对抗用户注入 |
| 多语言情感分析 | 普通 + 多语言模型 | 支持任何 HF 模型 |
| 模型路由(Query Classification) | 普通模式 | 按查询分配不同模型处理 |
| 产品评论分类 | 普通 + 批量处理 | Pipeline 批量预测 |
坑与注意
- 池化策略选错严重拉低准确率:对于 mean-pooling 模型(如 all-MiniLM-L6-v2)强制用 CLS 会导致语义相似度失效;建议默认 auto
- 新版保存的分类器有池化行为变化:0.2.0 前保存的模型需重新训练或显式传
pooling参数才能获得新行为 - 原型温度 0.25 默认值较保守:对类别边界清晰场景可尝试 0.1~0.15 提升区分度
- 少量新类别时神经头权重自动压制:如需固定比例(不推荐),同时设
new_class_prototype_weight和new_class_neural_weight - ONNX 加速仅 CPU:GPU 推理不走 ONNX,仍用 PyTorch;ONNX 主要对没有 GPU 的服务端有意义
- 幻觉检测精度有限:RAGTruth 基准 F1 仅 51.54%,生产环境应结合规则引擎使用,不宜单独依赖
- RAGTruth 基准数据偏特定场景:QA / Summarization / Data-to-Text 三类任务差异大(Summarization F1 仅 36.09%),选型前需在自己的数据集上验证
- head_steps 默认 300 步:低配机器上每轮约 1-2 秒;频繁增删示例时可降至 100 步以换取速度
与同类对比
| 维度 | Adaptive Classifier | HuggingFace Pipeline | sklearn + TF-IDF |
|---|---|---|---|
| 增量学习 | ✅ 不遗忘 | ❌ 需全量重训 | ❌ 需全量重训 |
| 动态增删类 | ✅ 运行时 | ❌ | ❌ |
| 对抗鲁棒性 | ✅ Strategic 模式 | ❌ | ❌ |
| ONNX 加速 | ✅ 内置 | ❌ | ❌ |
| 接入任意 HF 模型 | ✅ | ⚠️ 部分 | ❌ |
| 适合少样本场景 | ✅ 原型记忆 | ❌ | ⚠️ |
| 部署复杂度 | 中 | 低 | 低 |
Adaptive Classifier 填补了"动态增类 + 对抗鲁棒"这个空白,比传统 sklearn 方案灵活,比纯 LLM 调用方案轻量。
一句话推荐结论
需要分类器在生产环境持续学习新类别、且有对抗输入风险 → 选 Adaptive Classifier;只要简单固定类别分类 → sklearn + TF-IDF 更直接。