Skip to content

Scattering transform API documentation

ScatteringTransform

ScatteringTransform(
    W_adj, n_scales, n_layers, nlin=torch.abs, **kwargs
)

Bases: Module

ScatteringTransform base class. Inherits from PyTorch nn.Module.

This class implements the base logic to compute graph scattering transforms with a pooling and an arbitrary wavelet transform operators.

This is a base class, and implements only the logic to compute an arbitrary scattering transform. The method get_wavelets must be implemented by the subclass

Parameters:

Name Type Description Default
W_adj Tensor

Weighted adjacency matrix

required
n_scales int

Number of scales to use in wavelet transform

required
n_layers int

Number of layers in the scattering transform

required
nlin Callable[[Tensor], Tensor]

Non-linearity used in the scattering transform. Defaults to torch.abs

abs
**kwargs Any

Additional keyword arguments

{}
Source code in gsxform/scattering.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
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
def __init__(
    self,
    W_adj: torch.Tensor,
    n_scales: int,
    n_layers: int,
    nlin: Callable[[torch.Tensor], torch.Tensor] = torch.abs,
    **kwargs: Any,
) -> None:
    """Initialize scattering transform base class.

    This is a base class, and implements only the logic to compute
    an arbitrary scattering transform. The method `get_wavelets`
    must be implemented by the subclass

    Parameters
    ----------
    W_adj: torch.Tensor
        Weighted adjacency matrix
    n_scales: int
        Number of scales to use in wavelet transform
    n_layers: int
        Number of layers in the scattering transform
    nlin: Callable
        Non-linearity used in the scattering transform. Defaults to torch.abs
    **kwargs: Any
        Additional keyword arguments
    """
    super().__init__()

    # adjacency matrix, registered so .to(device) moves it with the module
    self.register_buffer("W_adj", W_adj)
    # number of scales
    self.n_scales = n_scales
    # number of layers
    self.n_layers = n_layers

    self.n_nodes = W_adj.shape[1]
    assert W_adj.shape[1] == W_adj.shape[2]

    self.nlin = nlin

forward

forward(x)

Forward pass of a generic scattering transform.

Parameters:

Name Type Description Default
x Tensor

input batch of graph signals

required

Returns:

Name Type Description
phi Tensor

scattering representation of the input batch

Source code in gsxform/scattering.py
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
def forward(self, x: torch.Tensor) -> torch.Tensor:
    """Forward pass of a generic scattering transform.

    Parameters
    ----------
    x: torch.Tensor
        input batch of graph signals

    Returns
    -------
    phi: torch.Tensor
        scattering representation of the input batch

    """
    batch_size = x.shape[0]

    n_features = x.shape[1]

    assert batch_size == self.W_adj.shape[0], (
        f"batch size of x ({batch_size}) must match W_adj ({self.W_adj.shape[0]})"
    )

    lowpass = self.get_lowpass(batch_size)
    psi = self.get_wavelets()

    # compute first scattering layer, low pass filter
    phi: torch.Tensor = einsum(x, lowpass, "b f n, b n -> b f")
    phi = rearrange(phi, "b f -> b f 1")

    # reshape inputs for loop
    S_x = rearrange(x, "b f n -> b 1 f n")

    for ll in range(1, self.n_layers):
        S_x_ll = torch.empty(
            [batch_size, 0, n_features, self.n_nodes], device=x.device
        )

        for jj in range(self.n_scales ** (ll - 1)):
            # intermediate repr, one copy per scale to contract against psi
            x_jj = repeat(S_x[:, jj, :, :], "b f n -> b ns f n", ns=self.n_scales)

            # wavelet filtering operation
            psi_x_jj = einsum(x_jj, psi, "b ns f n, b ns n m -> b ns f m")

            # application of non-linearity, yields scattering output
            S_x_jj = self.nlin(psi_x_jj)

            # concat scattering scale for the layer
            S_x_ll = torch.cat((S_x_ll, S_x_jj), dim=1)

            # compute scattering representation
            phi_jj = einsum(S_x_jj, lowpass, "b ns f n, b n -> b f ns")

            phi = torch.cat((phi, phi_jj), dim=2)

        S_x = S_x_ll.clone()  # continue iteration through the layer

    return phi

get_lowpass

get_lowpass(batch_size)

Compute lowpass filtering/pooling operator.

This should roughly resemble an average, it alters the output scaling factor. For instance averaging with the norm of the degree vector scales towards zero, this implementation offers a more natural scaling.

Parameters:

Name Type Description Default
batch_size int

Number of graphs in the batch

required

Returns:

Name Type Description
lowpass Tensor

average pooling operator

Source code in gsxform/scattering.py
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
def get_lowpass(self, batch_size: int) -> torch.Tensor:
    """Compute lowpass filtering/pooling operator.

    This should roughly resemble an average, it alters the output
    scaling factor. For instance averaging with the norm
    of the degree vector scales towards zero, this implementation
    offers a more natural scaling.

    Parameters
    ----------
    batch_size: int
        Number of graphs in the batch

    Returns
    -------
    lowpass: torch.Tensor
        average pooling operator
    """
    lowpass = (1 / self.n_nodes) * torch.ones(
        batch_size, self.n_nodes, device=self.W_adj.device
    )

    return lowpass

get_wavelets

get_wavelets()

Compute the wavelet operator.

Subclasses are required to implement this method.

Source code in gsxform/scattering.py
119
120
121
122
123
124
def get_wavelets(self) -> torch.Tensor:
    """Compute the wavelet operator.

    Subclasses are required to implement this method.
    """
    raise NotImplementedError

Diffusion

Diffusion(W_adj, n_scales, n_layers, nlin=torch.abs)

Bases: ScatteringTransform

Diffusion scattering transform.

Subclass of ScatteringTransform, implements get_wavelets method. Diffusion scattering transform algorithm based on description in Gama et. al 2018.

Parameters:

Name Type Description Default
W_adj Tensor

Weighted adjacency matrix

required
n_scales int

Number of scales to use in wavelet transform

required
n_layers int

Number of layers in the scattering transform

required
nlin Callable[[Tensor], Tensor]

Non-linearity used in the scattering transform. Defaults to torch.abs

abs
Source code in gsxform/scattering.py
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
def __init__(
    self,
    W_adj: torch.Tensor,
    n_scales: int,
    n_layers: int,
    nlin: Callable[[torch.Tensor], torch.Tensor] = torch.abs,
) -> None:
    """Initialize diffusion scattering transform.

    Parameters
    ----------
    W_adj: torch.Tensor
        Weighted adjacency matrix
    n_scales: int
        Number of scales to use in wavelet transform
    n_layers: int
        Number of layers in the scattering transform
    nlin: Callable[torch.Tensor]
        Non-linearity used in the scattering transform. Defaults to torch.abs

    """
    super().__init__(W_adj, n_scales, n_layers, nlin)

get_wavelets

get_wavelets()

Subclass method used to get wavelet filter bank.

This method returns diffusion wavelets

Returns:

Name Type Description
psi Tensor

diffusion wavelet operator

Source code in gsxform/scattering.py
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
def get_wavelets(self) -> torch.Tensor:
    """Subclass method used to get wavelet filter bank.

    This method returns diffusion wavelets

    Returns
    -------
    psi: torch.Tensor
        diffusion wavelet operator

    """
    # compute diffusion matrix
    T = lazy_diffusion(self.W_adj)
    # compute wavelet operator
    psi = diffusion_wavelets(T, self.n_scales)

    return psi

TightHann

TightHann(
    W_adj, n_scales, n_layers, nlin=torch.abs, use_warp=True
)

Bases: ScatteringTransform

TightHann scattering transform.

Subclass of ScatteringTransform, implements get_wavelets methods. Also additionally implements functions used to compute spectrum-adaptive wavelets.

Parameters:

Name Type Description Default
W_adj Tensor

Weighted adjacency matrix

required
n_scales int

Number of scales to use in wavelet transform

required
n_layers int

Number of layers in the scattering transform

required
nlin Callable[[Tensor], Tensor]

Non-linearity used in the scattering transform. Defaults to torch.abs

abs
use_warp bool

Use warping function. Defaults to True

True
Source code in gsxform/scattering.py
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
def __init__(
    self,
    W_adj: torch.Tensor,
    n_scales: int,
    n_layers: int,
    nlin: Callable[[torch.Tensor], torch.Tensor] = torch.abs,
    use_warp: bool = True,
) -> None:
    """Initialize tight Hann scattering transform.

    Parameters
    ----------
    W_adj: torch.Tensor
        Weighted adjacency matrix
    n_scales: int
        Number of scales to use in wavelet transform
    n_layers: int
        Number of layers in the scattering transform
    nlin: Callable[torch.Tensor]
        Non-linearity used in the scattering transform. Defaults to torch.abs
    use_warp: bool
        Use warping function. Defaults to True

    """
    super().__init__(W_adj, n_scales, n_layers, nlin)
    self.use_warp = use_warp
    self.warp = self.warp_func()

get_kernel

get_kernel()

Compute TightHann kernel adaptively.

Source code in gsxform/scattering.py
326
327
328
def get_kernel(self) -> TightHannKernel:
    """Compute TightHann kernel adaptively."""
    return TightHannKernel(self.n_scales, self.max_eig, self.warp)

get_wavelets

get_wavelets()

Subclass method used to get wavelet filter bank.

This method returns diffusion wavelets

Returns:

Name Type Description
psi Tensor

diffusion wavelet operator

Source code in gsxform/scattering.py
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
def get_wavelets(self) -> torch.Tensor:
    """Subclass method used to get wavelet filter bank.

    This method returns diffusion wavelets

    Returns
    -------
    psi: torch.Tensor
        diffusion wavelet operator

    """
    # compute wavelet operator
    psi = tighthann_wavelets(self.W_adj, self.n_scales, self.get_kernel())

    return psi

warp_func

warp_func()

Compute the spectrum-adaptive warping function.

Source code in gsxform/scattering.py
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
def warp_func(self) -> Callable[[torch.Tensor], torch.Tensor]:
    """Compute the spectrum-adaptive warping function."""
    E, V = compute_spectra(self.W_adj)
    # sort within each graph
    self.spectra, _ = torch.sort(E, dim=-1)
    self.max_eig = self.spectra[:, -1]

    n_eigs = self.spectra.shape[-1]
    cdf = torch.arange(
        0, n_eigs, device=self.spectra.device, dtype=self.spectra.dtype
    ) / (n_eigs - 1.0)
    cdf = cdf.expand_as(self.spectra)

    step = max(1, int(n_eigs / 5 - 1))

    if self.use_warp:
        xp, fp = self.spectra[:, 0::step], cdf[:, 0::step]
    else:
        xp, fp = self.spectra, cdf

    self.register_buffer("_warp_xp", xp.contiguous())
    self.register_buffer("_warp_fp", fp.contiguous())

    return lambda eig: _interp(eig, self._warp_xp, self._warp_fp)