山大软院创新项目实训个人博客——诈骗克星(三)
1. 本次进展
1)优化 train_wav2vec2.py,提升训练稳定性与不平衡数据适配能力
2)完成 Flask 服务封装 app.py,实现音频检测接口化
3)完成本地联调脚本 test_api.py,打通“上传音频 ->返回结构化结果”全流程
2. 训练脚本优化(train_wav2vec2.py)
2.1 优化背景
ASVspoof2019 LA 的类别分布不平衡,若仅用默认训练方式,可能偏向多数类,导致某类召回不足,因此将脚本升级为可调优版本,兼顾模型效果与工程可复现性
2.2 主要新增能力
(1) 类别不平衡:use_class_weight
1)启用后使用加权交叉熵
2)少数类权重更高,降低被忽视概率
示例代码:加权损失
class WeightedTrainer(Trainer):
def __init__(self, class_weights: Optional[torch.Tensor] = None, *args, **kwargs):
super().__init__(*args, **kwargs)
self.class_weights = class_weights
def compute_loss(self, model, inputs, return_outputs=False, **kwargs):
labels = inputs.pop("labels")
outputs = model(**inputs)
logits = outputs.logits
if self.class_weights is not None:
loss = F.cross_entropy(logits, labels, weight=self.class_weights.to(logits.device))
else:
loss = F.cross_entropy(logits, labels)
return (loss, outputs) if return_outputs else loss
(2) 训练稳定性参数
weight_decay:解耦权重衰减,每步都把参数往 0 拉一点,避免参数无限变大,抑制过拟合,提升泛化,模型更平滑
warmup_ratio:学习率预热比例(训练前期用较小学习率,逐渐升到设定值,减少训练初期震荡),具体原理为warmup_steps=warmup_ratio×total_steps,total为预先计算的总训练步数
grad_accum_steps:梯度累积步数
(显存/内存不够时,通过多步累积更新参数,等效增大 batch size,OOM或训练噪声大时调整)
save_total_limit:控制 checkpoint 数量(避免磁盘被大量模型快照占满,最好2或3)
(3) 阈值可调:decision_threshold
不再固定 0.5,可按反诈场景调召回/误报平衡
(4) 最优模型选择标准:metric_for_best_model
支持按 f1 / accuracy / f1_bonafide 从训练中每个epoch的模型中选出 best
3. Flask 服务化(app.py)
为了让后端可直接调用模型能力,将推理逻辑封装为 HTTP 服务。
3.1 接口设计
get("/"):服务说明(路由、模型、设备)
get("/health"):健康检查(加载模型状态)
post("/audio/detect"):上传音频并返回检测结果
3.2 返回结构(统一)
统一使用:code、message、request_id、data
/audio/detect 的 data 包含:
spoof/bonafide、spoof_prob、bonafide_prob、confidence、risk_level(low/medium/high)、latency_ms、model_version、device、sample_rate、max_seconds
3.3 工程细节
1)服务启动时按需加载模型(避免重复加载)
2)支持常见音频后缀校验
3)临时文件自动删除
4)标准化错误码(400x 参数问题,500x 服务异常)
3.4示例代码
@app.post("/audio/detect")
def audio_detect():
request_id = str(uuid.uuid4())
if "file" not in request.files:
return make_resp(4001, "缺少文件字段 file", http_status=400, request_id=request_id)
f = request.files["file"]
if not f.filename:
return make_resp(4002, "文件名为空", http_status=400, request_id=request_id)
suffix = Path(f.filename).suffix.lower()
if suffix not in {".wav", ".flac", ".mp3", ".m4a", ".ogg"}:
return make_resp(4003, f"不支持的音频格式: {suffix}", http_status=400, request_id=request_id)
save_path = UPLOAD_DIR / f"{int(time.time() * 1000)}_{Path(f.filename).name}"
try:
f.save(save_path)
result = predict_audio(save_path)
return make_resp(0, "success", result, request_id=request_id)
except Exception as e:
return make_resp(5001, str(e), http_status=500, request_id=request_id)
finally:
if save_path.exists():
try:
save_path.unlink()
except Exception:
pass
4. 本地联调脚本(test_api.py)
为了验证接口调用链路,我增加了 test_api.py:
接收参数:audio_path
自动向 /audio/detect 发起 POST 请求
打印 HTTP 状态码和 JSON 响应
PS D:\项目实训> python test_api.py --audio_path data/ASVspoof2019_LA_dev/flac/LA_D_1047731.flac
status_code: 200
{
"code": 0,
"data": {
"bonafide_prob": 0.9917102456092834,
"confidence": 0.9917102456092834,
"device": "cpu",
"label": "bonafide",
"latency_ms": 450.82,
"max_seconds": 4.0,
"model_version": "outputs\\wav2vec2-mid\\best",
"risk_level": "low",
"sample_rate": 16000,
"spoof_prob": 0.008289763703942299
},
"message": "success",
"request_id": "26a7a133-b626-45b0-913d-c46e29d66dcf"
}
作用:
1)验证接口可用性:直接请求 /audio/detect,确认服务是否真的跑起来了
2)验证请求格式:检查 multipart/form-data 上传文件这一层是否正确(字段名 file 是否匹配)
3)验证返回结构:看返回是否符合定义的协议:code / message / request_id / data,
以及 label/spoof_prob/... 等字段。
4)快速排障:服务未启动(连续拒绝)、路由问题(404)、参数问题(400x)、
模型推理异常(500x)
5.下一步计划
1)做阈值扫描实验(重点优化 spoof 漏检)
2)固化一版联调用接口协议文档(字段+错误码)
3)与后端进行正式联调,接入大模型工具调用链路
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)