1. 本次成果

将训练规模提升到中规模,验证模型在更多样本上的稳定性。

完成推理工具化:单条音频推理脚本 infer_audio.py、批量音频推理脚本 infer_audio_batch.py。

统一输出与训练脚本同类型的评估指标(spoof/bonafide 双视角 + 混淆矩阵)。

2. 中规模训练实践

2.1 训练脚本说明

本阶段仍使用统一训练脚本(已从 minimal 发展为通用训练脚本),核心流程:

(1)读取 ASVspoof2019 LA 的 CM protocol

(2)标签映射:bonafide -> 0, spoof -> 1

(3)加载 Wav2Vec2-Base 进行二分类微调

(4)输出整体指标,双视角指标与混淆矩阵

2.2 本次训练配置(中规模)

train 样本:2000       dev 样本:500

epoch:3(时间较长,至少跑了 1 轮得到有效结果)

每一轮训练可得一个checkpoint,并选出最优模型

2.3 指标结果

final eval metrics: {'eval_loss': 0.06812044233083725, 'eval_accuracy': 0.984, 'eval_f1': 0.9911111106089111, 'eval_precision_spoof': 0.9933184855211731, 'eval_recall_spoof': 0.9889135254966986, 'eval_f1_spoof': 0.9911111106089111, 'eval_precision_bonafide': 0.9019607842960401, 'eval_recall_bonafide': 0.938775510184923, 'eval_f1_bonafide': 0.9199999994818001, 'eval_cm_tp': 446, 'eval_cm_tn': 46, 'eval_cm_fp': 3, 'eval_cm_fn': 5, 'eval_pred_spoof_count': 449, 'eval_pred_bonafide_count': 51, 'eval_true_spoof_count': 451, 'eval_true_bonafide_count': 49, 'eval_runtime': 282.4445, 'eval_samples_per_second': 1.77, 'eval_steps_per_second': 0.443, 'epoch': 3.0}

各指标具体含义反诈克星(一)可见

2.4 结果解读

(1)对 spoof 的识别能力较强,召回与精确率都很高。

(2)bonafide 指标相对略低,原因是数据分布明显不均衡(spoof 数量远高于 bonafide)。

(3)FP/FN 都较低,说明模型具备较好的可用性基础。

3. 单条推理脚本 infer_audio.py

3.1 功能

对单条音频进行推理,输出结构化 JSON:

label:模型预测标签(spoof 或 bonafide)

spoof_prob:预测为 spoof 的概率

bonafide_prob:预测为 bonafide 的概率

confidence:预测类别对应概率(即 max(spoof_prob, bonafide_prob))

latency_ms:这条音频单次推理耗时(毫秒)

3.2 示例命令

python infer_audio.py --model_dir outputs/wav2vec2-mid/best --audio_path data/ASVspoof2019_LA_dev/flac/LA_D_1047731.flac

3.3 作用

该脚本是后续 Flask API 的核心推理逻辑基础,便于先做离线验证再做服务化封装。

4. 批量推理脚本 infer_audio_batch.py

4.1 功能

(1)批量读取目录下音频并推理

(2)结果逐条写入 CSV

(3)若提供 protocol(标明音频类型),可自动计算评估指标

4.2 CSV 字段(逐条)

audio_path:音频文件完整路径

utt_id:音频ID(文件名去掉后缀),用于和 protocol 对齐

label:模型预测标签(spoof 或 bonafide)

true_label:真实标签(来自 protocol;未提供 protocol 时为空)

spoof_prob:预测为 spoof 的概率

bonafide_prob:预测为 bonafide 的概率

confidence:预测类别对应概率(即 max(spoof_prob, bonafide_prob))

latency_ms:这条音频单次推理耗时(毫秒)

status:ok 或 error

error:异常信息(仅 status=error 时有内容)

4.3 示例命令(含评估)

python infer_audio_batch.py --model_dir outputs/wav2vec2-mid/best --audio_dir data/ASVspoof2019_LA_dev/flac --glob "*.flac" --protocol_path data/ASVspoof2019_LA_cm_protocols/ASVspoof2019.LA.cm.dev.trl.txt --output_csv outputs/dev_batch_results.csv

4.4关键代码

单条推理核心(概率、标签、置信度、耗时):

def predict_one(model, processor, device, audio_path: Path, sample_rate: int, max_seconds: float):
    wav = load_audio(audio_path, target_sr=sample_rate)
    max_len = int(max_seconds * sample_rate)
    if len(wav) > max_len:
        wav = wav[:max_len]

    inputs = processor(wav, sampling_rate=sample_rate, return_tensors="pt", padding=True)
    inputs = {k: v.to(device) for k, v in inputs.items()}

    start = time.time()
    with torch.no_grad():
        logits = model(**inputs).logits
        probs = torch.softmax(logits, dim=-1).cpu().numpy()[0]
    latency_ms = (time.time() - start) * 1000

    pred_id = int(np.argmax(probs))
    label = "bonafide" if pred_id == 0 else "spoof"

    return {
        "audio_path": str(audio_path),
        "utt_id": audio_path.stem,
        "label": label,
        "true_label": "",
        "spoof_prob": float(probs[1]),
        "bonafide_prob": float(probs[0]),
        "confidence": float(probs[pred_id]),
        "latency_ms": round(latency_ms, 2),
        "status": "ok",
        "error": "",
    }

批量扫描与容错(错误样本不中断全流程):

    rows = []
    start_all = time.time()
    for i, fp in enumerate(files, start=1):
        try:
            row = predict_one(model, processor, device, fp, args.sample_rate, args.max_seconds)
            if label_map:
                row["true_label"] = label_map.get(row["utt_id"], "")
        except Exception as e:
            row = {
                "audio_path": str(fp),
                "utt_id": fp.stem,
                "label": "",
                "true_label": "",
                "spoof_prob": "",
                "bonafide_prob": "",
                "confidence": "",
                "latency_ms": "",
                "status": "error",
                "error": str(e),
            }
        rows.append(row)
        if i % 50 == 0 or i == len(files):
            print(f"progress: {i}/{len(files)}")

4.5汇总输出

eval_accuracy(全局预测精确率)

eval_precision_spoof 、 eval_recall_spoof 、 eval_f1_spoof(以spoof为正类)

eval_precision_bonafide、 eval_recall_bonafide 、 eval_f1_bonafide(以bonafide为正类)

eval_cm_tp 、 eval_cm_tn 、 eval_cm_fp 、 eval_cm_fn(混淆矩阵)

预测分布与真实分布计数

5.小结

(1)将训练规模扩大至中等并顺利完成,输出可视化指标。但耗费时间过长。

(2)完成基于训练所得模型的音频推理脚本(单条/批量),为后续服务化和联调打下基础。

(3)接下来预计增加阈值调优与类别不均衡处理、封装 Flask 检测接口(/audio/detect)并与后端联调、用 eval 集做更完整测试。

Logo

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

更多推荐