目录

Audio Event Detection Model Training

基于深度可分离卷积 (DS-CNN) 的端侧音频事件检测模型训练工具,目标平台 openvela / TFLite Micro。

版本: v1.6.0 | 最佳准确率: 91.55% (4-class) | 6-class: 79.21%

快速开始

# 安装依赖
pip install tensorflow scipy scikit-learn matplotlib

# 4 类训练(初赛推荐)
python train_audio_event_model.py \
    --data-dir=../datasets_segments \
    --target-classes=knock,cough \
    --model-size=small --batch-size=128 --epochs=50 \
    --output-dir=output/small_4class \
    --spec-augment --noise-mix-prob=0.5 --delta-features

# 6 类训练(复赛,knock + cough + glass_breaking + dog_bark)
python train_audio_event_model.py \
    --data-dir=../datasets_segments \
    --target-classes=knock,cough,glass_breaking,dog_bark \
    --model-size=small --batch-size=128 --epochs=50 \
    --output-dir=output/small_6class \
    --spec-augment --noise-mix-prob=0.5 --delta-features

# PCEN 实验(替代 log-mel)
python train_audio_event_model.py \
    --data-dir=../datasets_segments \
    --target-classes=knock,cough \
    --model-size=small --batch-size=128 --epochs=50 \
    --output-dir=output/pcen \
    --spec-augment --noise-mix-prob=0.5 --delta-features \
    --pcen

数据集准备

1. 合并源数据

# 将 ESC-50、FSD50K、UrbanSound8K 按类别合并到 datasets/ 目录
# 每个类别一个子目录,包含所有 .wav 文件(可用软链接节省空间)

2. 峰值切分

# 使用 RMS 能量检测 + 事件聚类,切为 1 秒片段
python ../micro_speech/prepare_event_segments.py \
    --input_dir=datasets \
    --output_dir=datasets_segments \
    --target_classes=knock,cough,glass_breaking,dog_bark \
    --class_min_active_ratio=knock:0.03,cough:0.08,glass_breaking:0.05,dog_bark:0.04

3. 质量过滤

# 筛除低活动率、低 SNR 的劣质片段
python filter_segments.py \
    --active-ratio-min=0.05 --peak-rms-min=0.005 --snr-min-db=6

数据集结构

datasets_segments/
├── knock/              1015 segments
├── cough/              1043 segments
├── glass_breaking/     2669 segments   ← 新增
├── dog_bark/           3168 segments   ← 新增
├── background/         2496 segments
├── train/  (自动软链接, 8:1:1 预分割)
├── val/
└── test/

数据集来源

来源 使用类别 样本数 License
ESC-50 knock, cough, glass_breaking, dog_bark, background ~2,000 CC BY-NC
FSD50K knock, cough, glass_breaking, dog_bark, background ~4,000 CC
UrbanSound8K dog_bark, background ~3,000 CC
MUSAN background 930 CC / Public Domain

各类别数据量

类别 ESC-50 FSD50K UrbanSound8K 分段后 过滤后
knock 40 373 1,015 1,015
cough 40 385 1,043 1,043
glass_breaking 40 1,241 2,669 2,669
dog_bark 40 930 1,000 3,168 3,117
background 480 (12类环境声) Engine+Drill (1,619) 空调+引擎怠速+街景 (3,000) 2,496 2,496

当前背景音组成

来源 类别 数量
ESC-50 环境类 rain, wind, sea_waves, thunderstorm, crickets, chirping_birds, insects, water_drops, pouring_water, crackling_fire, vacuum_cleaner, washing_machine 480
background.wav 自定义录制 12 分钟环境音 144
其他 混合源

6 类数据分布

分集 knock cough glass dog bg silence 合计
train 811 835 2,135 2,493 1,996 auto ~8,270
val 102 104 267 312 250 auto ~1,035
test 102 104 267 312 250 auto ~1,035

关键参数

参数 默认值 说明
--data-dir 必填 数据集根目录
--target-classes 必填 目标事件类别,逗号分隔
--model-size small micro / tiny / small
--epochs 50 训练轮次
--batch-size 128 批次大小
--delta-features False 推荐。Mel + Δ + Δ²,+3% 准确率
--spec-augment False 推荐。频域/时域掩码增强
--noise-mix-prob 0.5 推荐。背景混音概率
--noise-snr-min-db 5 混音最小 SNR
--pcen False 实验性。PCEN 替代 log-mel
--pcen-alpha 0.5 PCEN AGC 强度,越大越压制背景
--pcen-delta 2.0 PCEN 偏置,越大越保留弱声
--pcen-root 0.25 PCEN 根压缩,越小压缩越强
--se-attention False 实验性。千参模型无提升
--focal-loss False 实验性。千参模型导致坍塌

工具脚本

脚本 用途
train_audio_event_model.py 主训练脚本
filter_segments.py 质量过滤:筛除低活动率/低 SNR 片段
../micro_speech/prepare_event_segments.py RMS 峰值检测 + 事件聚类切分原始音频

输出文件

文件 说明
model_int8.tflite int8 量化模型(MCU 部署)
model_float.tflite float32 模型(PC 端推理)
model.cc / model.h C 数组(TFLM 直接编译)
best_weights.weights.h5 最佳权重(可继续训练)
labels.txt 类别标签列表
metadata.json 模型参数与特征配置
metrics.json Keras float32 评估报告
metrics_tflite_int8.json int8 TFLite 评估报告
threshold_sweep.json 阈值 0.30-0.95 扫描 + 连续 N 帧触发
confusion_matrix.csv/png 混淆矩阵
training_curves.png 训练/验证曲线
roc_curves.png 各类别 ROC 曲线
feature_samples.png 各类别 mel 频谱样本

输出模型索引

输出目录 类别数 配置 准确率 TFLite
output/delta/ 4 Δ+Δ² + SpecAug + 混音 91.55% 12 KB
output/hard_noise/ 4 Mel + SNR 0-10dB 88.89% 12 KB
output/standard/ 4 Mel + SpecAug + 混音 88.41% 12 KB
output/multi4/ 6 Δ+Δ² + 混音 训练中

优化实验结论

方案 效果 结论
Δ+Δ² 特征 +3.1% (88.4% → 91.5%) 推荐
背景混音增强 降低漏检率 推荐
SpecAugment 提升泛化 推荐
SNR 0-10dB 训练 降低漏检,误报略升 ⚠️ 视场景选用
SE Attention -10.2% ❌ 不适合千参模型
Focal Loss 坍塌至 29.7% ❌ 不适合千参模型

模型架构

输入: 49 帧 × 40 Mel bins × C 通道 (C=3 启用 Δ)
    │
    ▼ Reshape to (49, 40, C)
Stem:   Conv2D(5×5, stride=2) → BN → ReLU
DS1:    DWConv(3×3) → BN → ReLU → PWConv(1×1) → BN → ReLU
DS2:    DWConv(3×3) → BN → ReLU → PWConv(1×1) → BN → ReLU
    │
    ▼ GlobalAveragePooling2D
    │
    ▼ Dense(num_classes, softmax)
输出: 各类别概率
型号 参数量 TFLite (int8)
micro ~700 ~6 KB
tiny ~960 ~10 KB
small (Mel) ~1548 ~12 KB
small (Δ+Δ²) ~1600 ~12 KB

特征预处理

采样率:      16 kHz
窗口:        30ms / 步长 20ms
FFT:         512 点 → 257 bins
Mel:         40 频带 (125-7500 Hz)
归一化:      log(mel+1e-6) → (+12)*1.625 → clip[0, 26]
Δ:           帧间差分 delta[t] = mel[t] - mel[t-1]
Δ²:          delta 的帧间差分
输出:        49 × 40 × C  (C=1 或 3)

### PCEN(可选, --pcen)

PCEN(mel) = (mel / (ε + smooth)^α + δ)^r − δ^r smooth[t] = 0.95·smooth[t-1] + 0.05·mel[t] (EMA, ~600ms时间常数)

α AGC增益控制,压制持续背景噪声 δ 偏置,保留微弱事件信号 r 根压缩,替代log的非线性变换

MCU代价: 40 floats (160 bytes) + ~200浮点运算/帧


## 更新日志

### v1.6.0 (2026-06-21)

- **新增** `--pcen` PCEN (Per-Channel Energy Normalization) 替代 log-mel
- **新增** `--pcen-alpha/delta/root` 可调 PCEN 参数
- PCEN 通过每频带 AGC 自适应抑制背景、放大弱声,适合噪声鲁棒场景
- 与 log-mel 同维度,不改模型架构

### v1.5.0 (2026-06-21)

- **新增** 6 分类支持:`knock` + `cough` + `glass_breaking` + `dog_bark`
- **新增** `filter_segments.py` 质量过滤脚本(active_ratio / peak_rms / SNR)
- **新增** 预分割目录自动检测(`train/val/test/` 结构优于随机切分)
- **新增** `threshold_sweep.json` 阈值扫描与连续触发模拟
- **新增** `metrics_tflite_int8.json` int8 TFLite 独立评估
- **修复** `stft(pad_end=False)` 帧数不一致导致 reshape 崩溃(零填充补齐)

### v1.4.0 (2026-06-21)

- **新增** `--se-attention`,实验结论:千参模型无提升(81.4%)

### v1.3.0 (2026-06-21)

- **新增** `--delta-features`,准确率 88.4% → 91.55% (+3.1%)

### v1.2.0 (2026-06-21)

- **新增** 4 种可视化图表(training_curves / confusion_matrix / roc / feature_samples)

### v1.1.0 (2026-06-21)

- **新增** `--focal-loss`,实验结论:千参模型不适用(29.7%)

### v1.0.0 (初始版本)

- DS-CNN + 背景混音 + SpecAugment + int8 量化 + C 数组导出
- 基线准确率 88.65%
关于

TFLite Micro,openvela

37.8 MB
邀请码
    Gitlink(确实开源)
  • 加入我们
  • 官网邮箱:gitlink@ccf.org.cn
  • QQ群
  • QQ群
  • 公众号
  • 公众号

版权所有:中国计算机学会技术支持:开源发展技术委员会
京ICP备13000930号-9 京公网安备 11010802047560号