這個專案實作了從時域 ConvTasNet 老師模型蒸餾到頻域 STFT-based 全連接學生模型的完整流程,用於**語音降噪(Speech Denoising)**任務。
新增功能: ✨ Quantization-Aware Training (QAT) - 支援 int8 量化,準備 DSP 部署!
kd/
├── models.py # Float32 學生模型(SubBandSTFTNet, SimpleSTFTNet)
├── losses.py # KD 損失函數
├── trainer.py # Two-step 訓練器
├── dataset.py # 合成降噪數據集
├── main.py # Float32 KD 主程式
│
├── quantization.py # ✨ Fake Quantization 模組
├── qat_models.py # ✨ QAT 學生模型
├── int_inference.py # ✨ Integer-only 推理引擎
├── main_qat.py # ✨ QAT 三階段訓練主程式
│
├── kd.md # KD 理論文件
├── kd-qat.md # QAT 理論文件
├── README.md # 本文件
├── README_QAT.md # ✨ QAT 完整指南
└── ARCHITECTURE.md # 系統架構視覺化
我們要解決的挑戰:
- 老師模型:ConvTasNet(時域、卷積+TCN、參數量大)
- 學生模型:只用 FC layers(頻域 STFT、參數量小)
- 目標:讓小模型學會大模型的能力
流程圖:
含噪語音 (mixture)
│
├─→ [老師:ConvTasNet] ─→ 增強語音 (time) ─→ STFT ─→ 老師頻譜
│ ↓
└─→ STFT ─→ [學生:FC-STFT] ─→ 增強頻譜 ───────────→ KD Loss
↓
+ mixture phase
↓
重建波形 ─────────────────────────→ Waveform KD Loss
Phase 1 (Pure KD): 學生純粹模仿老師
- L_mag_kd: STFT magnitude KD
- L_mask_kd: Mask KD(scale-invariant)
- L_wav_kd: Waveform KD
Phase 2 (Fine-tuning): 在真實乾淨語音上微調
- L_supervised: 與乾淨語音比較
- L_mag_kd (optional): 保留一點 KD 穩定訓練
pip install torch torchaudio numpy tqdm快速體驗完整的知識蒸餾流程:
python main.py --mode demo這會:
- 創建 ConvTasNet 老師和 STFT-FC 學生
- 生成少量合成數據
- 執行 Phase 1 + Phase 2 訓練(各 2-3 epochs)
- 展示推理結果
預期執行時間:5-10 分鐘(CPU)或 1-2 分鐘(GPU)
體驗完整的 Float32 → QAT → Integer 三階段訓練:
python main_qat.py --mode demo這會:
- Phase 1: Float32 KD 基準訓練
- Phase 2: QAT + KD(int8 量化感知訓練)
- Phase 3: Integer-only 驗證
預期執行時間:10-15 分鐘(CPU)或 2-4 分鐘(GPU)
結果:
- 壓縮比:54.3x(4.9M → 90K 參數)
- 精度:int8 量化,損失 ~23%
- 部署:準備好 DSP 部署
詳細文檔請見:README_QAT.md
使用更大的數據集和更多 epochs:
# Float32 KD 訓練
python main.py --mode train \
--phase1_epochs 10 \
--phase2_epochs 5 \
--batch_size 8 \
--device cuda
# QAT 訓練
python main_qat.py --mode train \
--phase1_epochs 10 \
--phase2_epochs 10 \
--device cuda參數說明:
--phase1_epochs: Phase 1 訓練輪數--phase2_epochs: Phase 2 訓練輪數--batch_size: Batch 大小--device: 'cuda' 或 'cpu'
使用訓練好的模型:
# Float32 推理
python main.py --mode inference \
--checkpoint checkpoints/best_phase2.pth \
--device cuda
# QAT 推理(支援 Float/Fake Quant/Integer 模式)
python main_qat.py --mode inference \
--checkpoint checkpoints_qat/phase2_qat_best.pth \
--device cuda頻譜分成多個 sub-bands,每個 band 用同一個小 MLP 處理
優點:
✓ 參數量小
✓ 保留頻率局部性
✓ 類似卷積的 inductive bias
每個時間 frame 獨立處理的 MLP
優點:
✓ 實作簡單
✓ 適合作為 baseline
- MagnitudeKDLoss: STFT magnitude 的 L1/MSE loss
- MaskKDLoss: Scale-invariant mask loss
- WaveformKDLoss: 時域波形的 L1/MSE loss
- MultiResolutionSTFTLoss: 多解析度 STFT loss
- SupervisedLoss: 與真實乾淨語音比較
- CombinedKDLoss: 組合以上 losses
KDTrainer 類別處理:
- STFT/iSTFT 轉換
- 老師模型推理
- 學生模型訓練
- Two-step training 邏輯
- Checkpoint 儲存/載入
SyntheticNoisyDataset:
- 合成「乾淨語音」(多諧波正弦波)
- 生成多種噪音(白噪音、粉紅噪音、布朗噪音)
- 控制 SNR(-5 ~ 20 dB)
- FakeQuantize: 核心量化層(支援 int8 symmetric/asymmetric)
- QuantizedLinear: Weight + Activation 雙量化的 FC 層
- FakeQuantSTFT: 模擬 DSP 的 Q15 STFT 量化
- QuantizedSubBandSTFTNet: QAT 版本的 SubBandSTFTNet
- QuantizedSimpleSTFTNet: QAT 版本的 SimpleSTFTNet
- 支援
enable_quantization()/disable_quantization() get_all_quantized_params(): 匯出 int8 參數
- IntegerLinearLayer: 純 int8 矩陣乘法
- IntegerReLU: 整數 ReLU
- IntegerSigmoid: LUT-based Sigmoid(簡化版)
- 為 DSP 部署準備的參考實作
- Phase 1: Float32 KD Baseline
- Phase 2: QAT + KD(int8 量化感知訓練)
- Phase 3: Integer-only 驗證
每個模組都可以獨立測試:
# Float32 模組
python models.py # 測試學生模型
python losses.py # 測試損失函數
python dataset.py # 測試數據集
python trainer.py # 測試訓練器
# QAT 模組
python quantization.py # ✨ 測試 Fake Quantization
python qat_models.py # ✨ 測試 QAT 模型
python int_inference.py # ✨ 測試 Integer 推理
python test_qat_minimal.py # ✨ 測試 QAT 訓練流程老師(ConvTasNet): ~4.9M 參數
學生(Float32): ~92K 參數 (壓縮 53.7x)
學生(QAT int8): ~91K 參數 (壓縮 54.3x)
在合成數據上:
- Phase 1 後:學生能模仿老師的大致行為
- Phase 2 後:學生進一步接近真實乾淨語音
- SNR 改善:通常 3-8 dB(取決於原始 SNR)
Phase 1 (Float32 KD): Val Loss ~90-100
Phase 2 (QAT + KD): Val Loss ~110-120 (+23.5%)
Phase 3 (驗證): Float vs Quant 差異 = 0.000000
量化精度:
- Weight: int8 [-128, 127]
- Activation: int8 [-128, 127]
- STFT: Q15 (可選)
部署優勢:
- 模型大小: ~0.3 MB (vs Float32 ~0.4 MB)
- 推理速度: ~2-3x 加速(在支援 int8 的硬體上)
- 記憶體: ~4x 減少
注意:實際性能取決於:
- 訓練數據的質量和數量
- 老師模型本身的性能
- 超參數設定
- 量化策略(symmetric/asymmetric, per-channel/per-tensor)
要在真實語音上使用,修改 dataset.py:
from torch.utils.data import Dataset
import torchaudio
class RealNoisyDataset(Dataset):
def __init__(self, clean_dir, noise_dir, ...):
# 載入真實的乾淨語音(如 LibriSpeech)
# 載入真實的噪音(如 DNS Challenge noise)
# 動態混合
pass推薦的真實數據集:
- 乾淨語音: LibriSpeech, VCTK
- 噪音: DNS Challenge, DEMAND, NOISEX-92
# 更小的模型(更快,但可能性能較差)
student = SubBandSTFTNet(
num_subbands=4,
hidden_dims=[64, 128, 64],
context_frames=1
)
# 更大的模型(性能更好,但較慢)
student = SubBandSTFTNet(
num_subbands=16,
hidden_dims=[256, 512, 256],
context_frames=3
)# Phase 1: 更重視波形 KD
loss_fn = CombinedKDLoss(
phase='phase1',
lambda_mag_kd=0.5,
lambda_mask_kd=0.5,
lambda_wav_kd=2.0 # 增加權重
)
# Phase 2: 純 supervised(不要 KD)
loss_fn = CombinedKDLoss(
phase='phase2',
mu_supervised=1.0,
mu_mag_kd=0.0 # 關閉 KD
)本實作基於以下研究:
- ConvTasNet: "Conv-TasNet: Surpassing Ideal Time-Frequency Magnitude Masking for Speech Separation"
- SubBand-KD: "Sub-Band Knowledge Distillation Framework for Speech Enhancement"
- Two-Step KD: "Two-Step Knowledge Distillation for Tiny Speech Enhancement"
- DISPatch: "Distilling Selective Patches for Speech Enhancement"
- DFKD: "Dynamic Frequency-Adaptive Knowledge Distillation for Speech Enhancement"
詳細理論請參考 kd.md。
Q: 為什麼學生只用 FC layers? A: 這是一個極端的壓縮場景,展示 KD 的能力。實務上可以用 Conv1D 或 Transformer。
Q: 為什麼不直接用老師模型? A: 老師太大太慢,不適合邊緣設備(如手機、嵌入式系統)。
Q: Phase 要用 mixture phase 嗎? A: 在降噪任務中,mixture phase 通常足夠好。也可以試試 oracle phase(真實乾淨語音的 phase)來看上限。
Q: 能用於其他任務嗎(如分離、增強)?
A: 可以!只需修改 ConvTasNet 的 num_sources 參數和數據集。
本專案僅供教學和研究使用。
感謝所有被引用論文的作者們的貢獻!