| Home | Publications | Projects | News | CV | GitHub · Google Scholar · Email |
Trainable multilinear bases for image compression and inpainting, learned on the unitary manifold in JAX
pdft is a JAX library for learning parametric quantum Fourier transforms by manifold optimization. It replaces the fixed DFT or DCT with a basis that is trained per dataset, while keeping near-linear transform cost, exact invertibility, and a parameter count polylogarithmic in the image size.
The core idea: optimize circuit parameters on Riemannian manifolds (products of U(2) and U(1) groups) to learn frequency-domain representations that are more compressible than the standard DFT. On Quick Draw line drawings, the trained basis stores images in roughly 20% fewer bytes than JPEG’s 8×8 block cosine transform at the same reconstruction quality.
The same circuits, with the Hadamards frozen, are the trainable transforms of the follow-up paper on image inpainting (arXiv:2609.17298) — see Coherence and image inpainting below.
pdft is the maintained implementation. It supersedes the original Julia package ParametricDFT.jl, which is slated for archival — see Julia lineage below.

Figure 2 from the paper — four circuit variants and the DCT-IV acting on an input image x. A box spanning two legs is a general two-qubit tensor U(4) ∈ U(4), shared by (a)–(c), which differ only in wiring. A bond with two endpoint dots is a controlled phase M ∈ U(1)4, used in (d) both within and across the two wires.
The variants keep the same local gate and differ in how those tensors are wired:
Optimization runs on the manifold of unitary matrices, so every iterate stays exactly orthonormal and the learned transform remains invertible by construction. Two Riemannian optimizers are provided: gradient descent and Adam.
import jax
import jax.numpy as jnp
import pdft
target = jax.random.normal(jax.random.PRNGKey(7), (4, 4)).astype(jnp.complex128)
basis = pdft.QFTBasis(m=2, n=2)
result = pdft.train_basis(
basis,
target=target,
loss=pdft.L1Norm(),
optimizer=pdft.RiemannianGD(lr=0.01),
steps=50,
seed=0,
)
Because the whole stack is JAX, training is JIT-compiled and runs on CPU, GPU, or TPU without a separate accelerator backend.
Beyond training, the package covers JSON and compression I/O, visualization, and runnable demos under examples/ for basis training, optimizer benchmarking, and MERA.
Compression cares only how few coefficients a basis needs. Recovering an image from a subset of its pixels — inpainting, completion, compressed sensing — is governed by a second quantity, the coherence of the basis with the pixel basis,
| μ(U) = N maxij | Uij | 2 ∈ [1, N], |
| and the number of samples needed for recovery grows linearly with it. Every basis in pdft starts at μ = 1, and there is a structural reason it can stay there: if the only non-diagonal gates are one Hadamard per wire, then | Uij | = N−1/2 for every parameter value, so μ = 1 identically. Training the controlled-phase gates arbitrarily hard, on any objective, cannot move it. Training the Hadamard or U(4) gates can and does. |
The pdft.coherence module makes this checkable before a run rather than measured after it: certify_flat_modulus reports whether a basis with a given set of frozen gates keeps μ = 1 identically, and otherwise returns exactly which gates must be held fixed, in the form that train_basis_batched accepts.
The paper Quantum-Inspired Trainable and Parameter-Efficient Tensor Networks for Image Inpainting (arXiv:2609.17298, with Konstantinos Slavakis, submitted to ICASSP 2027) builds on exactly this. With the Hadamards frozen the circuit is a diagonal relaxation of the QFT: unitary for every parameter value, O(N2 log N) to apply to an N × N image, and trainable with plain Adam through an unrolled hard-thresholding recovery, with no Riemannian retractions and no coherence penalty. On DIV2K at 10% observed pixels, the model with 288 trainable phases outperforms the DFT by 2.0 dB and the best fixed transform by 1.6 dB in PSNR, and comes within 0.2 dB of a learned butterfly factorization with 64 times as many parameters. Freeing the Hadamards as well buys a further 0.2 dB, but coherence then drifts above 1 and the retractions return.
pdft began as a port of ParametricDFT.jl and is a feature-complete one: parity against the Julia reference is verified by committed golden files covering training routines, I/O, and visualization. With the port complete, the Julia package is being retired in favor of pdft.
One piece of that work outlived the port. GPU-accelerated Riemannian optimization, originally written to bypass limitations in existing Julia manifold libraries, was upstreamed into ManifoldsGPU.jl and now serves the JuliaManifolds ecosystem independently of this project.
| Package | Role |
|---|---|
| JAX | Autodiff, JIT compilation, accelerators |
| NumPy | Array manipulation and I/O |
| Matplotlib | Visualization of bases and reconstructions |
Last updated September 22, 2026.