69 lines
2.5 KiB
Python
69 lines
2.5 KiB
Python
"""可选:用 transformers 训练 BERT 级意图分类器(需 requirements-ml.txt)。
|
|
|
|
这是"正式环境"路线(对照实现方案 3.2 路由器选型表):
|
|
- 规则分类器(当前默认): 90-93% 准确率,零成本,适合 MVP
|
|
- 小 LLM / BERT 分类器: 94-97% 准确率,适合正式环境
|
|
|
|
用法(示例):
|
|
python scripts/train_classifier.py --data data/train.jsonl --output models/classifier
|
|
|
|
数据格式(每行一个 JSON):
|
|
{"query": "...", "domain": "code|math|legal|medical|general"}
|
|
|
|
注意:本脚本是流水线骨架,需自行准备数据。
|
|
"""
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import json
|
|
import sys
|
|
from pathlib import Path
|
|
|
|
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
|
|
|
|
|
def main():
|
|
parser = argparse.ArgumentParser()
|
|
parser.add_argument("--data", required=True, help="训练数据 jsonl 路径")
|
|
parser.add_argument("--output", default="models/classifier", help="输出目录")
|
|
parser.add_argument("--base", default="Qwen/Qwen3-0.6B", help="基础模型")
|
|
parser.add_argument("--epochs", type=int, default=3)
|
|
args = parser.parse_args()
|
|
|
|
try:
|
|
from transformers import (AutoModelForSequenceClassification, AutoTokenizer,
|
|
TrainingArguments)
|
|
except ImportError as e:
|
|
print("需要 ML 依赖:pip install -r requirements-ml.txt")
|
|
raise SystemExit(1) from e
|
|
|
|
# 数据格式转换
|
|
samples = []
|
|
with open(args.data, encoding="utf-8") as f:
|
|
for line in f:
|
|
line = line.strip()
|
|
if line:
|
|
obj = json.loads(line)
|
|
samples.append(obj)
|
|
print(f"加载 {len(samples)} 条训练样本")
|
|
|
|
labels = ["code", "math", "legal", "medical", "general"]
|
|
label2id = {l: i for i, l in enumerate(labels)}
|
|
|
|
# 这里仅演示训练流水线;正式训练请使用 datasets 库构建 Dataset。
|
|
tokenizer = AutoTokenizer.from_pretrained(args.base)
|
|
model = AutoModelForSequenceClassification.from_pretrained(args.base, num_labels=5)
|
|
|
|
print("训练参数示例:")
|
|
print(TrainingArguments(
|
|
output_dir=args.output,
|
|
num_train_epochs=args.epochs,
|
|
per_device_train_batch_size=8,
|
|
learning_rate=2e-5,
|
|
))
|
|
print("请参考实现方案 3.2 路由器选型表,构建带标注数据集后执行正式训练。")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main()
|
|
|