Skip to content

Graph kernel functions API documentation

TightHannKernel

TightHannKernel(n_scales, max_eig, omega=None)

TightHannKernel class.

Thie class constructs a spectrum-adaptive tight-hann kernel function used in its corresponding wavelet transform. Based off of the implementation from Tabar et. al 2021 of the algorithm originally described in Shuman et. al 2015

Parameters:

Name Type Description Default
n_scales int

number of scales used in wavelet transform

required
max_eig Tensor

the maximum eigenvalue of the graph laplacian. Used for scaling purposes

required
omega Callable[[Tensor], Tensor] | None

warping function. Defaults to None

None
Source code in gsxform/kernel.py
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
def __init__(
    self,
    n_scales: int,
    max_eig: torch.Tensor,
    omega: Callable[[torch.Tensor], torch.Tensor] | None = None,
) -> None:
    """Initialize TightHannKernel class.

    Parameters
    ----------
    n_scales: int
        number of scales used in wavelet transform
    max_eig: torch.Tensor:
        the maximum eigenvalue of the graph laplacian. Used for scaling purposes
    omega: Union[Callable[[torch.Tensor], torch.Tensor], None]
        warping function. Defaults to None

    """
    self.n_scales = n_scales
    self.K = 1
    self.R = 3.0
    self.max_eig = max_eig

    if omega is not None:
        self.omega = omega
        self.max_eig = self.omega(self.max_eig.float())
    else:
        self.omega = lambda eig: eig

    if self.n_scales + 1 - self.R <= 0:
        raise ValueError(
            f"n_scales must be greater than {self.R - 1:g}, got {n_scales}"
        )

    # dilation factor, might need to reverse this to account for swapped bounds...
    # self.d = (self.M + 1 - self.R) / (self.R * self.max_eig)
    self.d = self.R * self.max_eig / (self.n_scales + 1 - self.R)

    # add a trailing axis so it broadcasts against the (batch, n_nodes)
    # eigenvalue tensors
    if self.d.ndim > 0:
        self.d = self.d.unsqueeze(-1)
    # hann kernel functional form
    self.kernel: Callable[[torch.Tensor], torch.Tensor] = lambda eig: (
        sum(
            [
                0.5 * torch.cos(2 * np.pi * (eig / self.d - 0.5) * k)
                for k in range(self.K + 1)
            ]
        )
        * (eig >= 0)
        * (eig <= self.d)
    )

get_adapted_kernel

get_adapted_kernel(eig, scale)

Compute spectrum adapted kernels.

return self.kernel(self.omega(eig) - self.d / self.R * (scale - self.R + 1))

Parameters:

Name Type Description Default
eig Tensor

input tensor of eigenvalues of the graph laplacian

required
scale int

The scale parameter of the specific kernel. Not to be confused with n_scales, which is the total number of scales used by the wavelet transform.

required

Returns:

Name Type Description
adapted_kernel Tensor

scale-specific adapted kernel

Source code in gsxform/kernel.py
 78
 79
 80
 81
 82
 83
 84
 85
 86
 87
 88
 89
 90
 91
 92
 93
 94
 95
 96
 97
 98
 99
100
101
def get_adapted_kernel(self, eig: torch.Tensor, scale: int) -> torch.Tensor:
    """Compute spectrum adapted kernels.

    return self.kernel(self.omega(eig) - self.d / self.R * (scale - self.R + 1))

    Parameters
    ----------
    eig: torch.Tensor
        input tensor of eigenvalues of the graph laplacian
    scale: int
        The scale parameter of the specific kernel. Not to be confused
        with `n_scales`, which is the total number of scales used
        by the wavelet transform.

    Returns
    -------
    adapted_kernel: torch.Tensor
        scale-specific adapted kernel

    """
    adapted_kernel = self.kernel(
        self.omega(eig) - self.d / self.R * (scale - self.R + 1)
    )
    return adapted_kernel