Filter basis functions#

To parameterize spherical discrete-continuous (DISCO) convolutions, we need to pick suitable basis functions. Currently, torch-harmonics supports 3 different basis types:

  • Picewise Linear basis functions (Similar to hat functions in finite element methods)

  • Morlet-wavelet-like basis functions defined on the sphere

  • Zernike polynomials

This notebook outlines depicts the different types of basis functions.

import math

import numpy as np
import torch

import matplotlib.pyplot as plt
cmap = plt.cm.RdBu

from torch_harmonics.filter_basis import get_filter_basis

Piecewise linear basis on a disk#

We start off with piecewise linear basis functions defined on a disk.

fb = get_filter_basis((5, 4), "piecewise linear")
phi = torch.linspace(0, 2*np.pi, 100)
r = torch.linspace(0, 1.0, 40)
r, phi = torch.meshgrid(r, phi)

# Quadrature on the unit disk: area element r dr dφ. Trapezoidal in r, uniform in φ.
nr, nphi = r.shape[0], r.shape[1]
dr = 1.0 / (nr - 1) if nr > 1 else 1.0
dphi = 2 * math.pi / nphi
r_fac = torch.ones_like(r)
r_fac[0, :] = 0.5
r_fac[-1, :] = 0.5
quad_weights = r * dr * dphi * r_fac

ks = fb.kernel_size
idx, vals = fb.compute_support_vals(r, phi, r_cutoff=1)
idx = torch.stack([idx[:, 0], idx[:, 1], idx[:, 2]], dim=0)
psi = torch.sparse_coo_tensor(idx, vals, size=(ks, *r.shape)).to_dense()

fig, ax = plt.subplots(nrows=1, ncols=ks, subplot_kw={"projection": "polar"}, figsize=(32,4))
for k in range(0, ks):
    f = psi[k]
    pc = ax[k].contourf(phi, r, f, cmap=cmap, vmin=-f.abs().max(), vmax=f.abs().max(), levels=24, antialiased=True)
    ax[k].set_xticks([])
    ax[k].set_yticks([])
/home/bbonev/.local/lib/python3.10/site-packages/torch/functional.py:554: UserWarning: torch.meshgrid: in an upcoming release, it will be required to pass the indexing argument. (Triggered internally at /pytorch/aten/src/ATen/native/TensorShape.cpp:4314.)
  return _VF.meshgrid(tensors, **kwargs)  # type: ignore[attr-defined]
../_images/b1ec844ad2a451a9220e2fbef942161b44c89e6b77e6a17e9f9f96dee02c8c57.png

adding up all basis functions to obtain a filter

filter = (fb.get_init_factors().reshape(-1, 1, 1) * psi).sum(dim=0)

fig = plt.figure()
ax = fig.add_subplot(projection="polar")
ax.contourf(phi, r, filter, cmap=cmap, vmin=-filter.abs().max(), vmax=filter.abs().max(), levels=24, antialiased=True)
ax.set_xticks([])
ax.set_yticks([])
plt.show()
../_images/38856b67305a571a451b856dea67e6e89ae6d48db612cd7558b75cc6026f943a.png
plt.plot(r[:, 0], filter[:, 0])
[<matplotlib.lines.Line2D at 0x70ce7c16fe50>]
../_images/9f63af4f85d53caf0c7bd825ce036b28bbd422d3146260ebb947bf267ef35481.png

Wavelet style basis#

This basis uses a Hann windowing function and a Fourier basis.

nmax = mmax = 3
fb = get_filter_basis((nmax, mmax), "harmonic")
phi = torch.linspace(0, 2*np.pi, 100)
r = torch.linspace(0, 1.0, 40)
r, phi = torch.meshgrid(r, phi)

ks = fb.kernel_size
idx, vals = fb.compute_support_vals(r, phi, r_cutoff=1)
idx = torch.stack([idx[:, 0], idx[:, 1], idx[:, 2]], dim=0)
psi = torch.sparse_coo_tensor(idx, vals, size=(ks, *r.shape)).to_dense()

fig, ax = plt.subplots(nrows=1, ncols=ks, subplot_kw={"projection": "polar"}, figsize=(32,4))
for k in range(0, ks):
    f = psi[k]
    pc = ax[k].contourf(phi, r, f, cmap=cmap, vmin=-f.abs().max(), vmax=f.abs().max(), levels=24, antialiased=True)
    ax[k].set_xticks([])
    ax[k].set_yticks([])
../_images/1949ff5a3616c37487b79599e3984b06f9fcc1c6a1aa61881156dca0610838d6.png
filter = (fb.get_init_factors().reshape(-1, 1, 1) * psi).sum(dim=0)

fig = plt.figure()
ax = fig.add_subplot(projection="polar")
ax.contourf(phi, r, filter, cmap=cmap, vmin=-filter.abs().max(), vmax=filter.abs().max(), levels=24, antialiased=True)
ax.set_xticks([])
ax.set_yticks([])
plt.show()
../_images/ef761f3b81aa51e02cd2a43ac832f55af81d81d28bf855066459125cf17cf8d3.png
dr = r[1, 0] - r[0, 0]
dphi = phi[0, 1] - phi[0, 0]
norms_sq = (psi**2 * r.unsqueeze(0) * dr * dphi).sum(dim=(-2, -1))
print("L2 norms (should be ~1):", norms_sq.sqrt())
L2 norms (should be ~1): tensor([1.0046, 1.0000, 1.0077, 1.0081, 1.0000, 1.0132, 1.0023, 1.0000, 1.0039])

Zernike polynomials#

for more info read https://en.wikipedia.org/wiki/Zernike_polynomials

nmax = 4
fb = get_filter_basis(nmax, "zernike")
phi = torch.linspace(0, 2*np.pi, 100)
r = torch.linspace(0, 1.0, 40)
r, phi = torch.meshgrid(r, phi)

ks = fb.kernel_size
idx, vals = fb.compute_support_vals(r, phi, r_cutoff=1)
idx = torch.stack([idx[:, 0], idx[:, 1], idx[:, 2]], dim=0)
psi = torch.sparse_coo_tensor(idx, vals, size=(ks, *r.shape)).to_dense()

fig, ax = plt.subplots(nrows=1, ncols=ks, subplot_kw={"projection": "polar"}, figsize=(32,4))
for k in range(0, ks):
    f = psi[k]
    pc = ax[k].contourf(phi, r, f, cmap=cmap, vmin=-f.abs().max(), vmax=f.abs().max(), levels=24, antialiased=True)
    ax[k].set_xticks([])
    ax[k].set_yticks([])
../_images/083faacdcffdfd7ebcd6aff2d610dedbe91e6de18327297185fba62f0be5a6f5.png

add abasis function to obtain a filter

filter = (fb.get_init_factors().reshape(-1, 1, 1) * psi).sum(dim=0)

fig = plt.figure()
ax = fig.add_subplot(projection="polar")
ax.contourf(phi, r, filter, cmap=cmap, vmin=-filter.abs().max(), vmax=filter.abs().max(), levels=24, antialiased=True)
ax.set_xticks([])
ax.set_yticks([])
plt.show()
../_images/5c1a2db99263eea70c0e01246098d10f17cdab57e40332e7f38804a296224c93.png
plt.plot(r[:, 0], filter[:, 0])
[<matplotlib.lines.Line2D at 0x70ce7e340ee0>]
../_images/f54527d189ae45b3b9ec5cdde3ce57a339aafc390ee453d9c984ec2b7bd87bd2.png
dr = r[1, 0] - r[0, 0]
dphi = phi[0, 1] - phi[0, 0]
norms_sq = (psi**2 * r.unsqueeze(0) * dr * dphi).sum(dim=(-2, -1))
print("L2 norms (should be ~1):", norms_sq.sqrt())
L2 norms (should be ~1): tensor([1.0178, 1.0256, 1.0359, 1.0385, 1.0442, 1.0490, 1.0515, 1.0532, 1.0637,
        1.0621])

Fourier-Bessel Basis#

The Fourier-Bessel Basis functions on a disk are defined as

\[ \Psi_{m, n} (r, \theta) = J_m(\alpha_{m, n} r) \cos(m \theta) \]

Unlike, the Zernike polynomials, the Fourier-Bessel basis functiosn vanish smoothly to \(0\) at the boundary, which is advantageous for filtering and leakage purposes.

kernel_shape = (4,3)
fb = get_filter_basis(kernel_shape=kernel_shape, basis_type="fourier-bessel")
print(fb.kernel_size)
20
phi = torch.linspace(0, 2*np.pi, 100)
r = torch.linspace(0, 1.0, 40)
r, phi = torch.meshgrid(r, phi)

ks = fb.kernel_size
idx, vals = fb.compute_support_vals(r, phi, r_cutoff=1)
idx = torch.stack([idx[:, 0], idx[:, 1], idx[:, 2]], dim=0)
psi = torch.sparse_coo_tensor(idx, vals, size=(ks, *r.shape)).to_dense()

fig, ax = plt.subplots(nrows=1, ncols=ks, subplot_kw={"projection": "polar"}, figsize=(32,4))
for k in range(0, ks):
    f = psi[k]
    pc = ax[k].contourf(phi, r, f, cmap=cmap, vmin=-f.abs().max(), vmax=f.abs().max(), levels=24, antialiased=True)
    ax[k].set_xticks([])
    ax[k].set_yticks([])
../_images/9601fc292da36b1b866739b0607e018062e60986012129695d77511131f0e085.png
filter = (fb.get_init_factors().reshape(-1, 1, 1) * psi).sum(dim=0)

fig = plt.figure()
ax = fig.add_subplot(projection="polar")
ax.contourf(phi, r, filter, cmap=cmap, vmin=-filter.abs().max(), vmax=filter.abs().max(), levels=24, antialiased=True)
ax.set_xticks([])
ax.set_yticks([])
plt.show()
../_images/5b3eab04d7ff0eb3e9a9231f30d4c09f00b11f327a713caac13fa9785d41704e.png
plt.plot(r[:, 0], filter[:, 0])
[<matplotlib.lines.Line2D at 0x70ce76283eb0>]
../_images/e2e819862843c7f5381884700bfc9a3baa2fe21ab66e25f3d2c05eeb8965b99b.png
# R = 1.0
# phi = torch.linspace(0, 2*np.pi, 100)
# r = torch.linspace(0, R, 40)
# r, phi = torch.meshgrid(r, phi)

# # Quadrature on the unit disk: area element r dr dφ. Trapezoidal in r, uniform in φ.
# nr, nphi = r.shape[0], r.shape[1]
# dr = R / (nr - 1) if nr > 1 else 1.0
# dphi = 2 * math.pi / nphi
# r_fac = torch.ones_like(r)
# r_fac[0, :] = 0.5
# r_fac[-1, :] = 0.5
# quad_weights = r * dr * dphi * r_fac
# (quad_weights * psi[3].abs().pow(2)).sum().sqrt()
dr = r[1, 0] - r[0, 0]
dphi = phi[0, 1] - phi[0, 0]
norms_sq = (psi**2 * r.unsqueeze(0) * dr * dphi).sum(dim=(-2, -1))
print("L2 norms (should be ~1):", norms_sq.sqrt())
L2 norms (should be ~1): tensor([1.0048, 1.0101, 1.0000, 1.0101, 1.0000, 1.0046, 1.0101, 1.0000, 1.0101,
        1.0000, 1.0101, 1.0000, 1.0043, 1.0101, 1.0000, 1.0101, 1.0000, 1.0101,
        1.0000, 1.0040])