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)与后端进行正式联调,接入大模型工具调用链路

Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐