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 批量预测

坑与注意

  1. 池化策略选错严重拉低准确率:对于 mean-pooling 模型(如 all-MiniLM-L6-v2)强制用 CLS 会导致语义相似度失效;建议默认 auto
  2. 新版保存的分类器有池化行为变化:0.2.0 前保存的模型需重新训练或显式传 pooling 参数才能获得新行为
  3. 原型温度 0.25 默认值较保守:对类别边界清晰场景可尝试 0.1~0.15 提升区分度
  4. 少量新类别时神经头权重自动压制:如需固定比例(不推荐),同时设 new_class_prototype_weight 和 new_class_neural_weight
  5. ONNX 加速仅 CPU:GPU 推理不走 ONNX,仍用 PyTorch;ONNX 主要对没有 GPU 的服务端有意义
  6. 幻觉检测精度有限:RAGTruth 基准 F1 仅 51.54%,生产环境应结合规则引擎使用,不宜单独依赖
  7. RAGTruth 基准数据偏特定场景:QA / Summarization / Data-to-Text 三类任务差异大(Summarization F1 仅 36.09%),选型前需在自己的数据集上验证
  8. 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 更直接。