feat: 多专业小模型+路由模型系统 MVP(mock 全链路 + FastAPI 网关 + 论文调研)
This commit is contained in:
@@ -0,0 +1,68 @@
|
||||
"""可选:用 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()␍
|
||||
Reference in New Issue
Block a user