feat: 添加多 GPU 支持,优化训练脚本以支持全局批量大小
基于深度可分离卷积 (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
# 将 ESC-50、FSD50K、UrbanSound8K 按类别合并到 datasets/ 目录 # 每个类别一个子目录,包含所有 .wav 文件(可用软链接节省空间)
# 使用 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
# 筛除低活动率、低 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/
--data-dir
--target-classes
--model-size
small
micro
tiny
--epochs
--batch-size
--delta-features
--spec-augment
--noise-mix-prob
--noise-snr-min-db
--pcen
--pcen-alpha
--pcen-delta
--pcen-root
--se-attention
--focal-loss
train_audio_event_model.py
filter_segments.py
../micro_speech/prepare_event_segments.py
model_int8.tflite
model_float.tflite
model.cc
model.h
best_weights.weights.h5
labels.txt
metadata.json
metrics.json
metrics_tflite_int8.json
threshold_sweep.json
confusion_matrix.csv/png
training_curves.png
roc_curves.png
feature_samples.png
output/delta/
output/hard_noise/
output/standard/
output/multi4/
输入: 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) 输出: 各类别概率
采样率: 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
版权所有:中国计算机学会技术支持:开源发展技术委员会 京ICP备13000930号-9 京公网安备 11010802047560号
Audio Event Detection Model Training
基于深度可分离卷积 (DS-CNN) 的端侧音频事件检测模型训练工具,目标平台 openvela / TFLite Micro。
版本: v1.6.0 | 最佳准确率: 91.55% (4-class) | 6-class: 79.21%
快速开始
数据集准备
1. 合并源数据
2. 峰值切分
3. 质量过滤
数据集结构
数据集来源
各类别数据量
当前背景音组成
6 类数据分布
关键参数
--data-dir--target-classes--model-sizesmallmicro/tiny/small--epochs--batch-size--delta-features--spec-augment--noise-mix-prob--noise-snr-min-db--pcen--pcen-alpha--pcen-delta--pcen-root--se-attention--focal-loss工具脚本
train_audio_event_model.pyfilter_segments.py../micro_speech/prepare_event_segments.py输出文件
model_int8.tflitemodel_float.tflitemodel.cc/model.hbest_weights.weights.h5labels.txtmetadata.jsonmetrics.jsonmetrics_tflite_int8.jsonthreshold_sweep.jsonconfusion_matrix.csv/pngtraining_curves.pngroc_curves.pngfeature_samples.png输出模型索引
output/delta/output/hard_noise/output/standard/output/multi4/优化实验结论
模型架构
特征预处理
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浮点运算/帧