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.")