Skip to content

Repository files navigation

知識蒸餾:ConvTasNet → STFT-FC Network(含 QAT)

這個專案實作了從時域 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        # 系統架構視覺化

🎯 核心概念

問題設定

我們要解決的挑戰:

  1. 老師模型:ConvTasNet(時域、卷積+TCN、參數量大)
  2. 學生模型:只用 FC layers(頻域 STFT、參數量小)
  3. 目標:讓小模型學會大模型的能力

解決方案

流程圖:

含噪語音 (mixture)
    │
    ├─→ [老師:ConvTasNet] ─→ 增強語音 (time) ─→ STFT ─→ 老師頻譜
    │                                                        ↓
    └─→ STFT ─→ [學生:FC-STFT] ─→ 增強頻譜 ───────────→ KD Loss
                      ↓
                 + mixture phase
                      ↓
                  重建波形 ─────────────────────────→ Waveform KD Loss

Two-Step Training

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

1. Float32 KD Demo(基礎)

快速體驗完整的知識蒸餾流程:

python main.py --mode demo

這會:

  • 創建 ConvTasNet 老師和 STFT-FC 學生
  • 生成少量合成數據
  • 執行 Phase 1 + Phase 2 訓練(各 2-3 epochs)
  • 展示推理結果

預期執行時間:5-10 分鐘(CPU)或 1-2 分鐘(GPU)

2. ✨ QAT Demo(進階 - 支援 int8 量化)

體驗完整的 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

3. 完整訓練

使用更大的數據集和更多 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'

4. 推理

使用訓練好的模型:

# 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

📊 模組說明

models.py - 學生模型

SubBandSTFTNet(推薦)

頻譜分成多個 sub-bands,每個 band 用同一個小 MLP 處理

優點:
  ✓ 參數量小
  ✓ 保留頻率局部性
  ✓ 類似卷積的 inductive bias

SimpleSTFTNet

每個時間 frame 獨立處理的 MLP

優點:
  ✓ 實作簡單
  ✓ 適合作為 baseline

losses.py - 損失函數

  • MagnitudeKDLoss: STFT magnitude 的 L1/MSE loss
  • MaskKDLoss: Scale-invariant mask loss
  • WaveformKDLoss: 時域波形的 L1/MSE loss
  • MultiResolutionSTFTLoss: 多解析度 STFT loss
  • SupervisedLoss: 與真實乾淨語音比較
  • CombinedKDLoss: 組合以上 losses

trainer.py - 訓練器

KDTrainer 類別處理:

  1. STFT/iSTFT 轉換
  2. 老師模型推理
  3. 學生模型訓練
  4. Two-step training 邏輯
  5. Checkpoint 儲存/載入

dataset.py - 數據集

SyntheticNoisyDataset

  • 合成「乾淨語音」(多諧波正弦波)
  • 生成多種噪音(白噪音、粉紅噪音、布朗噪音)
  • 控制 SNR(-5 ~ 20 dB)

✨ QAT 相關模組

quantization.py - Fake Quantization

  • FakeQuantize: 核心量化層(支援 int8 symmetric/asymmetric)
  • QuantizedLinear: Weight + Activation 雙量化的 FC 層
  • FakeQuantSTFT: 模擬 DSP 的 Q15 STFT 量化

qat_models.py - QAT 學生模型

  • QuantizedSubBandSTFTNet: QAT 版本的 SubBandSTFTNet
  • QuantizedSimpleSTFTNet: QAT 版本的 SimpleSTFTNet
  • 支援 enable_quantization() / disable_quantization()
  • get_all_quantized_params(): 匯出 int8 參數

int_inference.py - Integer 推理引擎

  • IntegerLinearLayer: 純 int8 矩陣乘法
  • IntegerReLU: 整數 ReLU
  • IntegerSigmoid: LUT-based Sigmoid(簡化版)
  • 為 DSP 部署準備的參考實作

main_qat.py - QAT 主程式

  • 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)

Float32 KD 性能

在合成數據上:

  • Phase 1 後:學生能模仿老師的大致行為
  • Phase 2 後:學生進一步接近真實乾淨語音
  • SNR 改善:通常 3-8 dB(取決於原始 SNR)

✨ QAT 性能

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 減少

注意:實際性能取決於:

  1. 訓練數據的質量和數量
  2. 老師模型本身的性能
  3. 超參數設定
  4. 量化策略(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
)

調整 KD 權重

# 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
)

📚 參考文獻

本實作基於以下研究:

  1. ConvTasNet: "Conv-TasNet: Surpassing Ideal Time-Frequency Magnitude Masking for Speech Separation"
  2. SubBand-KD: "Sub-Band Knowledge Distillation Framework for Speech Enhancement"
  3. Two-Step KD: "Two-Step Knowledge Distillation for Tiny Speech Enhancement"
  4. DISPatch: "Distilling Selective Patches for Speech Enhancement"
  5. 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 參數和數據集。

📝 授權

本專案僅供教學和研究使用。

🙏 致謝

感謝所有被引用論文的作者們的貢獻!

About

Knowledge Distillation + Quantization-Aware Training: ConvTasNet → STFT-FC Network for Speech Denoising (54.3x compression, int8 quantization, DSP-ready)

Resources

Stars

2 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages