Audio Codec Comparison: Discrete vs Continuous Bottlenecks
By Fouzil Ali · Published July 21, 2026
Compares three bottleneck strategies (discrete RVQ, VAE with KL, plain AE) on ESC-50 audio using DAC-style encoder/decoder architecture with spectral and perceptual loss metrics.
- audio-codec
- bottleneck-comparison
- rvq
- vae
- esc-50
15 cells7 experiments19 views0 forks
Inside this notebook
# Audio Codec Comparison: Discrete vs Continuous Bottleneck Comparing three bottleneck strategies on the **ESC-50** dataset: | Variant | Bottleneck | Key Losses | |---|---|---| | **Discrete (DAC-style)** | RVQ, 4×256 codebooks, dim=8 | MSE + Commitment (0.25) + Codebook (1.0) + Spectral | | **VAE w/ KL** | Gaussian VAE, β=0.1 | MSE + β·KL + Spectral | | **AE w/o KL** | Plain continuous | MSE + Spectral | **Uses a simplified DAC encoder/decoder:** WNConv1d + Snake1d, 4 stages with strides [2,4,8,8], bottleneck dim 512.
# ── 0. Install Dependencies ──
import subprocess, sys, importlib
deps = ['librosa', 'soundfile', 'datasets', 'matplotlib', 'tqdm', 'torchcodec', 'ipywidgets']
for pkg in deps:
try:
importlib.import_module(pkg)
print(f" ✓ {pkg}")
except ImportError:
print(f" Installing {pkg}...")
subprocess.check_call([sys.executable, '-m', 'pip', 'install', '-q', pkg])
print(f" ✓ {pkg} installed")
print("\nAll deps ready.")✓ librosa ✓ soundfile ✓ datasets ✓ matplotlib ✓ tqdm Installing torchcodec... ✓ torchcodec installed Installing ipywidgets... ✓ ipywidgets installed All deps ready.
# ── 1. Imports & ESC-50 Download (GitHub, extract via soundfile) ──
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader, random_split
import torchaudio
import matplotlib.pyplot as plt
from tqdm.auto import tqdm
import os, urllib.request, zipfile
from io import BytesIO
from pathlib import Path
import soundfile as sf
import warnings
warnings.filterwarnings('ignore')
DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
…Using device: cuda Downloading ESC-50 from GitHub (~650 MB)... Download complete. Extracting audio files... Found 2000 WAV files Using first 500 files Saved audio tensor: torch.Size([500, 1, 80000]) Audio tensor: torch.Size([500, 1, 80000]) Range: [-1.0000, 1.0000] Train: 400, Val: 100 Data ready!
# Inspect the zip structure
import zipfile
p = "/home/user/esc50_data/esc50_master.zip"
with zipfile.ZipFile(p, 'r') as zf:
names = zf.namelist()
print(f"Total entries: {len(names)}")
# Show first 30
for n in names[:30]:
print(f" {n}")
# Find audio files
audio = [n for n in names if n.endswith('.wav')]
print(f"\nWAV files: {len(audio)}")
if audio:
print(f" First: {audio[0]}")
print(f" Last: {audio[-1]}")
else:
# Check for any audio extension
audio_exts = [n for n in names if any(n.endswith(e) for e in ['.wav','.mp3','.flac','.ogg'])]
…Total entries: 2017 ESC-50-master/ ESC-50-master/.circleci/ ESC-50-master/.circleci/config.yml ESC-50-master/.github/ ESC-50-master/.github/stale.yml ESC-50-master/.gitignore ESC-50-master/LICENSE ESC-50-master/README.md ESC-50-master/audio/ ESC-50-master/audio/1-100032-A-0.wav ESC-50-master/audio/1-100038-A-14.wav ESC-50-master/audio/1-100210-A-36.wav ESC-50-master/audio/1-100210-B-36.wav ESC-50-master/audio/1-101296-A-19.wav ESC-50-master/audio/1-101296-B-19.wav…
# ── 2. Model Definitions: Encoder, Decoder, Bottlenecks ──
import math
class Snake1d(nn.Module):
"""Periodic activation used in DAC: x + (1/α) * sin²(αx)"""
def __init__(self, channels):
super().__init__()
self.alpha = nn.Parameter(torch.ones(1, channels, 1))
def forward(self, x):
return x + (1.0 / (self.alpha + 1e-9)) * torch.sin(self.alpha * x).pow(2)
class WNConv1d(nn.Module):
"""Weight-normalized Conv1d"""
def __init__(self, in_ch, out_ch, kernel_size, stride=1, padding=0, dilation=1):
super().__init__()
self.conv = nn.utils.weight_norm(
…Sanity: input torch.Size([2, 1, 16000]) → output torch.Size([2, 1, 16000]) VAE aux keys: dict_keys(['kl']) AE: input torch.Size([2, 1, 16000]) → output torch.Size([2, 1, 16000]) All model classes defined and verified.
# ── 3. Loss Functions, Metrics & Training Loop ──
def multi_scale_spectral_loss(est, ref, n_ffts=[512, 1024, 2048]):
"""Multi-scale spectral convergence loss (L1 on log-magnitude STFT)."""
loss = 0.0
est = est.squeeze(1) # [B, T]
ref = ref.squeeze(1)
for n_fft in n_ffts:
hop = n_fft // 4
win = torch.hann_window(n_fft).to(est.device)
est_stft = torch.stft(est, n_fft, hop, window=win, return_complex=True)
ref_stft = torch.stft(ref, n_fft, hop, window=win, return_complex=True)
est_mag = torch.abs(est_stft) + 1e-8
ref_mag = torch.abs(ref_stft) + 1e-8
loss += F.l1_loss(torch.log(est_mag), torch.log(ref_mag))
return loss / len(n_ffts)
def si_snr(est, ref):
…Training function defined. Ready to run variants.
# Quick test: verify train_variant is in kernel state and works for 1 epoch
print("Testing train_variant on main branch...")
result = train_variant('ae_nokl', 'Test-AE', n_epochs=1)
print(f"Test OK: {result['label']}, val_loss={result['best_val_loss']:.4f}")
print("Kernel state is ready for batch experiment.")