Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
45 changes: 45 additions & 0 deletions CHANGELOG
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,51 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0

### Added

- **Explicit `.compile()` / `.freeze()` for filters and chains.** Coefficient
computation stays lazy by default; `FX.compile(fs)` (and `FilterChain.compile`)
now eagerly designs coefficients ahead of `forward()` so the module no longer
mutates on its first call — useful for deployment, reproducibility, and as a
prerequisite for graph capture (`torch.compile` works, graph-breaking around the
native kernel). `freeze(fs=None)` additionally locks coefficients (skips the
per-forward fs/fingerprint recompute). Works on a lone deterministic filter
(`LoButterworth(...).compile(48000)`) as well as a chain. `FilterChain.summary()`
renders the topology and the fused SOS plan. (Note: full `torch.export` of the
forward still requires registering the native SOS kernel as a custom op — a
follow-up; the kernel is currently opaque to the exporter.)
- **`torchfx.ddsp` — opt-in differentiable filters.** The differentiable cascade
(`ddsp.differentiable_sos_cascade`, `ddsp.BiquadFunction`) wraps the native forward
kernel in a `torch.autograd.Function` with a hand-derived analytic backward (the
adjoint of an LTI filter reuses the same forward kernel), avoiding the ~30× cost of
unrolling the IIR recursion under autograd. The deterministic filters are unchanged
(still `@torch.no_grad()`), so non-DDSP users pay no overhead.
- **Extension point:** `ddsp.LearnableFilter` is an abstract base that mirrors the
deterministic `filter.AbstractFilter` — subclass it, declare design params as
`nn.Parameter`, and implement `sos()` (returning a differentiable `[K, 6]` stack)
to get a differentiable `forward` and a `freeze()` bridge for free.
- **Ready-made subclasses:** `LearnableLowpass`, `LearnableHighpass`,
`LearnablePeaking`, and `LearnableParametricEQ`. `freeze()` bakes trained params
into a deterministic `filter.SOSFilter` for fast inference.
- **Learnable effects:** `ddsp.LearnableGain` shows that pointwise effects are
trainable with plain autograd (`freeze()` → `effect.Gain`). Recursive/native-kernel
effects (Delay, Reverb, dynamics) instead each need their own analytic-backward
`autograd.Function` (the `BiquadFunction` pattern) and are not covered automatically.
- The `examples/ddsp_graphic_eq.py` PoC now uses this package (time-domain cascade)
instead of its frequency-domain workaround. References: RBJ cookbook, torchlpc,
"Differentiable All-Pole Filters".
- **`filter.SOSFilter`.** A deterministic filter wrapping a precomputed `[K, 6]`
SOS matrix (the inference-side target of `LearnableFilter.freeze()`).
- **CLI: SoX-style positional pipeline and a `compile` artifact.** `torchfx pipe
IN OUT lowpass --cutoff 800 reverb --mix 0.4` applies a positional effect
pipeline; `torchfx compile "<pipeline>" --fs 48000 -o chain.fxg` freezes a
pipeline into a portable `.fxg` artifact (filters stored as precomputed SOS,
other effects as re-instantiable specs), and `torchfx process --compiled
chain.fxg` applies it without re-designing coefficients. New `parse_pipeline`
parser reuses the existing effect registry.
- **Realtime smoothed parameter automation.** `RealtimeProcessor.set_parameter`
gained a `ramp_ms` argument: a numeric parameter is linearly ramped to its target
over that many milliseconds (at block granularity) instead of jumping, removing
zipper noise. Filter coefficients are recomputed each block during a ramp without
resetting DF1 state, so cutoff/Q sweeps stay continuous.
- **PipeWire realtime backend (#55).** `torchfx.realtime.PipeWireBackend` targets
PortAudio's PipeWire host (falling back to its PulseAudio host, which PipeWire
also serves), reusing the `SoundDeviceBackend` callback path — native low-latency
Expand Down
188 changes: 37 additions & 151 deletions examples/ddsp_graphic_eq.py
Original file line number Diff line number Diff line change
@@ -1,28 +1,25 @@
#!/usr/bin/env python3
"""Differentiable DSP proof-of-concept: a learnable parametric EQ in TorchFX.
"""Differentiable DSP: a learnable parametric EQ built on ``torchfx.ddsp``.

This example demonstrates that TorchFX's SOS filter representation is usable in a
fully differentiable, gradient-trained pipeline. A cascade of RBJ peaking-EQ
biquads — the same second-order sections the native TorchFX kernels execute — is
parameterised by learnable per-band ``(center frequency, Q, gain)`` and applied in
the **frequency domain** (FFT x H), where plain autograd differentiates everything:
no custom backward pass, no modification of the fast ``@torch.no_grad()`` inference
kernels.
This example trains a cascade of RBJ peaking-EQ biquads — the *same* second-order
sections the native TorchFX kernels execute — by gradient descent. Filtering happens
in the **time domain** through :func:`torchfx.ddsp.differentiable_sos_cascade`, which
runs the fast native forward kernel and a hand-derived analytic backward (no autograd
unrolling of the IIR recursion, and the deterministic ``@torch.no_grad()`` filters are
left untouched).

Task: invert an unknown coloration filter. A synthetic "speaker" response (a few
fixed resonances + a low shelf) colors a noise signal; the EQ is trained with a
multi-resolution STFT loss so that ``eq(colored)`` matches the original signal.
Gradients flow loss -> STFT -> waveform -> H(e^jw) -> RBJ coefficients ->
(fc, Q, gain) end to end.
fixed resonances) colors a noise signal; the EQ is trained with a multi-resolution
STFT loss so that ``eq(colored)`` matches the original signal. Gradients flow
``loss -> STFT -> waveform -> SOS cascade -> RBJ coefficients -> (fc, Q, gain)``.

Run (CPU is fine; CUDA used when available)::

python examples/ddsp_graphic_eq.py # train + save figure/audio
python examples/ddsp_graphic_eq.py # train + save figure
python examples/ddsp_graphic_eq.py --steps 0 # just plot the untrained EQ

The trained band parameters can be transferred 1:1 onto native TorchFX biquads for
fast no-grad inference, since both share the RBJ SOS parameterisation.

After training, ``eq.freeze()`` bakes the learned coefficients into a deterministic
``torchfx.filter.SOSFilter`` for fast no-grad inference.
"""

from __future__ import annotations
Expand All @@ -31,130 +28,15 @@
import math

import torch
from torch import Tensor, nn

# --------------------------------------------------------------------------- #
# Differentiable RBJ peaking-EQ cascade (frequency-domain application). #
# --------------------------------------------------------------------------- #


def rbj_peaking_sos(fc: Tensor, q: Tensor, gain_db: Tensor, fs: float) -> Tensor:
"""RBJ-cookbook peaking-EQ section as a differentiable ``[K, 6]`` SOS stack.

Identical math to ``torchfx.filter.biquad.BiquadPeak`` but in batched torch ops
so the coefficients stay on the autograd tape: gradients flow back to ``fc``,
``q`` and ``gain_db``.

"""
A = torch.pow(10.0, gain_db / 40.0)
w0 = 2.0 * math.pi * fc / fs
cos_w0 = torch.cos(w0)
alpha = torch.sin(w0) / (2.0 * q)

b0 = 1.0 + alpha * A
b1 = -2.0 * cos_w0
b2 = 1.0 - alpha * A
a0 = 1.0 + alpha / A
a1 = -2.0 * cos_w0
a2 = 1.0 - alpha / A

sos = torch.stack([b0 / a0, b1 / a0, b2 / a0, torch.ones_like(a0), a1 / a0, a2 / a0], dim=-1)
return sos


def sos_freq_response(sos: Tensor, n_fft: int) -> Tensor:
"""Complex frequency response of an SOS cascade on the ``rfft`` grid.

``H(z) = prod_k (b0 + b1 z^-1 + b2 z^-2) / (1 + a1 z^-1 + a2 z^-2)`` evaluated at
``z = e^{jw}`` for the ``n_fft//2 + 1`` non-negative frequency bins. All ops are
complex-differentiable, so this is the autograd path from coefficients to audio.

"""
w = torch.linspace(0.0, math.pi, n_fft // 2 + 1, device=sos.device, dtype=sos.dtype)
z1 = torch.exp(torch.complex(torch.zeros_like(w), -w)) # z^-1
z2 = z1 * z1
b = sos[:, 0:1] + sos[:, 1:2] * z1 + sos[:, 2:3] * z2 # [K, F]
a = sos[:, 3:4] + sos[:, 4:5] * z1 + sos[:, 5:6] * z2
return torch.prod(b / a, dim=0) # [F]


class LearnableParametricEQ(nn.Module):
"""A cascade of peaking-EQ bands with learnable (fc, Q, gain) per band.

Raw parameters are unconstrained; sigmoid/exp maps (FLAMO-style) keep the
effective values in stable, audible ranges so training cannot push a pole
outside the unit circle.

"""

def __init__(
self,
n_bands: int = 10,
fs: float = 48_000,
f_lo: float = 40.0,
f_hi: float = 16_000.0,
max_gain_db: float = 18.0,
) -> None:
super().__init__()
self.fs = fs
self.f_lo, self.f_hi = f_lo, f_hi
self.max_gain_db = max_gain_db
# Initialise band centers log-spaced over [f_lo, f_hi]; the sigmoid map
# below is centered so raw zeros land back on this grid.
grid = torch.linspace(0.0, 1.0, n_bands + 2)[1:-1]
self._fc_raw = nn.Parameter(torch.logit(grid))
self._q_raw = nn.Parameter(torch.zeros(n_bands)) # exp map: q0 * e^raw
self.gain_db_raw = nn.Parameter(torch.zeros(n_bands))

@property
def fc(self) -> Tensor:
"""Band centers, log-spaced sigmoid map into [f_lo, f_hi]."""
t = torch.sigmoid(self._fc_raw)
return self.f_lo * (self.f_hi / self.f_lo) ** t

@property
def q(self) -> Tensor:
"""Band Q, exp map around 1.0, clamped to [0.3, 8]."""
return torch.clamp(math.sqrt(2.0) * torch.exp(self._q_raw), 0.3, 8.0)

@property
def gain_db(self) -> Tensor:
"""Band gain in dB, tanh-bounded to +/- max_gain_db."""
return self.max_gain_db * torch.tanh(self.gain_db_raw)

def sos(self) -> Tensor:
"""The cascade as a differentiable ``[K, 6]`` SOS stack (TorchFX layout)."""
return rbj_peaking_sos(self.fc, self.q, self.gain_db, self.fs)

def forward(self, x: Tensor) -> Tensor:
"""Filter ``x`` (``[C, T]``) through the cascade in the frequency domain."""
n_fft = 2 ** math.ceil(math.log2(x.shape[-1]))
H = sos_freq_response(self.sos().to(torch.float64), n_fft).to(torch.complex64)
X = torch.fft.rfft(x, n=n_fft)
y = torch.fft.irfft(X * H, n=n_fft)
return y[..., : x.shape[-1]]


# --------------------------------------------------------------------------- #
# Multi-resolution STFT loss (the standard DDSP reconstruction objective). #
# --------------------------------------------------------------------------- #


def multires_stft_loss(x: Tensor, y: Tensor, sizes: tuple[int, ...] = (512, 1024, 2048)) -> Tensor:
"""Sum of spectral-convergence + log-magnitude L1 over several STFT resolutions."""
loss = x.new_zeros(())
for n_fft in sizes:
win = torch.hann_window(n_fft, device=x.device)
X = torch.stft(x, n_fft, n_fft // 4, window=win, return_complex=True).abs()
Y = torch.stft(y, n_fft, n_fft // 4, window=win, return_complex=True).abs()
loss = loss + torch.norm(X - Y) / (torch.norm(Y) + 1e-8)
loss = loss + torch.nn.functional.l1_loss(torch.log(X + 1e-5), torch.log(Y + 1e-5))
return loss


# --------------------------------------------------------------------------- #
# The unknown coloration to invert (fixed, non-learnable). #
# --------------------------------------------------------------------------- #
from torch import Tensor

from torchfx.ddsp import (
LearnableParametricEQ,
multires_stft_loss,
rbj_peaking_sos,
sos_freq_response,
)
from torchfx.filter import SOSFilter


def make_coloration(fs: float) -> Tensor:
Expand All @@ -179,19 +61,19 @@ def main() -> None:
torch.manual_seed(0)

# Training signal: pink-ish noise (white noise shaped by 1/sqrt(f)).
T = int(args.fs * args.seconds)
n_fft = 2 ** math.ceil(math.log2(T))
white = torch.randn(1, T, device=device)
n_samples = int(args.fs * args.seconds)
n_fft = 2 ** math.ceil(math.log2(n_samples))
white = torch.randn(1, n_samples, device=device)
freqs = torch.fft.rfftfreq(n_fft, 1 / args.fs, device=device)
shape = 1.0 / torch.sqrt(torch.clamp(freqs, min=20.0))
x = torch.fft.irfft(torch.fft.rfft(white, n=n_fft) * shape, n=n_fft)[..., :T]
x = torch.fft.irfft(torch.fft.rfft(white, n=n_fft) * shape, n=n_fft)[..., :n_samples]
x = x / x.abs().max()

# Color the signal with the unknown response (no grad — it is the plant).
# Color the signal with the unknown response (deterministic native path, no grad).
color_sos = make_coloration(args.fs).to(device)
colorizer = SOSFilter(color_sos, fs=args.fs).to(device)
with torch.no_grad():
Hc = sos_freq_response(color_sos.to(torch.float64), n_fft).to(torch.complex64)
colored = torch.fft.irfft(torch.fft.rfft(x, n=n_fft) * Hc, n=n_fft)[..., :T]
colored = colorizer(x)

# The learnable EQ must invert the coloration: eq(colored) ~= x.
eq = LearnableParametricEQ(n_bands=args.bands, fs=args.fs).to(device)
Expand All @@ -210,11 +92,11 @@ def main() -> None:

# ----- report ----------------------------------------------------------- #
with torch.no_grad():
H_eq = sos_freq_response(eq.sos().to(torch.float64), n_fft)
H_col = sos_freq_response(color_sos.to(torch.float64), n_fft)
h_eq = sos_freq_response(eq.sos().to(torch.float64), n_fft)
h_col = sos_freq_response(color_sos.to(torch.float64), n_fft)
f = torch.fft.rfftfreq(n_fft, 1 / args.fs)
eq_db = 20 * torch.log10(H_eq.abs() + 1e-9).cpu()
col_db = 20 * torch.log10(H_col.abs() + 1e-9).cpu()
eq_db = 20 * torch.log10(h_eq.abs() + 1e-9).cpu()
col_db = 20 * torch.log10(h_col.abs() + 1e-9).cpu()
residual_db = eq_db + col_db # perfect inversion -> 0 dB everywhere

band = (f > 60) & (f < 12_000)
Expand All @@ -224,6 +106,10 @@ def main() -> None:
for fc_i, q_i, g_i in zip(eq.fc.cpu(), eq.q.cpu(), eq.gain_db.cpu(), strict=True):
print(f" {fc_i:8.1f} {q_i:5.2f} {g_i:+6.2f}")

# The trained EQ can be frozen into a fast deterministic filter for inference.
frozen = eq.freeze()
print(f"frozen to {type(frozen).__name__} for fast no-grad inference")

try:
import matplotlib

Expand Down
11 changes: 10 additions & 1 deletion src/cli/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,8 @@

import typer

from cli.commands import info, play, preset, process, record, sox, watch
from cli.commands import compile as compile_mod
from cli.commands import info, pipeline, play, preset, process, record, sox, watch
from cli.repl import interactive_cmd

# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -68,6 +69,14 @@ def get_state() -> dict[str, object]:
# ---------------------------------------------------------------------------

app.command(name="process", help="Apply effects to audio files.")(process.process_cmd)
app.command(
name="pipe",
help="Apply a SoX-style positional effect pipeline to a file.",
context_settings={"allow_extra_args": True, "ignore_unknown_options": True},
)(pipeline.pipe_cmd)
app.command(name="compile", help="Freeze an effect pipeline into a portable .fxg artifact.")(
compile_mod.compile_cmd
)
app.command(name="info", help="Display audio file metadata.")(info.info_cmd)
app.command(name="play", help="Play an audio file through speakers.")(play.play_cmd)
app.command(name="record", help="Record audio from a microphone.")(record.record_cmd)
Expand Down
Loading
Loading