Run this notebook yourself!
Download the executed notebook: steerable_pyramid.ipynb!
Steerable Pyramids#
Warning
This notebook requires the optional dependency plenoptic, which can be installed with pip.
In this tutorial, we will show how fenestration can be used on dictionaries of 4d tensors, like the output of steerable pyramids. Here, we will use the SteerablePyramidFreq object from plenoptic.
The steerable pyramid can be thought of as a bank of oriented bandpass convolutional filters which span all orientations and frequencies. In this way, it is often thought of as having a representation similar to that of the primary visual cortex (V1). For an introduction to steerable pyramids, we recommend plenoptic’s user guide and the tutorial from pyrtools.
import fenestration as fen
import plenoptic as po
%load_ext autoreload
%autoreload 2
%matplotlib inline
Setting Up Figures#
First, we will load in an example image. These must be 4d tensors with shape batch by channel by height by width, which is the default for plenoptic’s images. We will also take advantage of plenoptic’s plotting functions in this tutorial.
img = po.data.einstein()
po.plot.imshow(img);
Creating the Steerable Pyramid#
Now we will initialize the SteerablePyramidFreq object and extract its coefficients for this image. The parameter height refers to the height of the pyramid (i.e., the number of spatial scales), which should be the same as the scales parameter we pass to PoolingWindows. Here we see that pyr_coeffs is a dictionary with keys 0 to 3, corresponding to height=4 in addition to residual_lowpass and residual_highpass. The shape of each of the dictionary elements now has an additional dimension corresponding to each of the orientations (four orientations is the default).
# create the pyramid
pyr = po.process.SteerablePyramidFreq(img.shape[-2:], height=4)
# get the pyramid coefficients; this is equivalent to pyr.forward(img)
pyr_coeffs = pyr(img)
print(pyr_coeffs.keys())
print(pyr_coeffs[0].shape)
odict_keys(['residual_highpass', 0, 1, 2, 3, 'residual_lowpass'])
torch.Size([1, 1, 4, 256, 256])
We can visualize these coefficient outputs using po.plot.pyrshow. Each plot shows the coefficients for a given scale and orientation band. Height 00 shows the finest scales with high spatial frequencies and we can see the spatial frequencies decrease as we move down the rows. Each column shows a given orientation (“band”), with vertical in the first column, horizontal in the third, and the diagonals in the second and fourth. Finally, the plots at the very bottom of the figure show the residuals, the highest and lowest frequencies, respectively, which are not captured by the pyramid.
po.plot.pyrshow(pyr_coeffs);
Applying PoolingWindows to Steerable Pyramid#
Now we can pass the pyramid coefficients into PoolingWindows in order to get the pooled windows at each scale and orientation. However, the input to PoolingWindow’s forward method must be a dictionary of 4d tensors (or a single 4d tensor) in which the keys are (scale, orientation) tuples. Therefore, we remove residual_highpass and residual_lowpass and rearrange these values into a new dictionary new_pyr.
new_pyr = {}
for k, v in pyr_coeffs.items():
# skip residuals
if isinstance(k, str):
continue
for ori in range(v.shape[2]):
new_pyr[(k, ori)] = v[:, :, ori, :, :]
print(new_pyr.keys())
dict_keys([(0, 0), (0, 1), (0, 2), (0, 3), (1, 0), (1, 1), (1, 2), (1, 3), (2, 0), (2, 1), (2, 2), (2, 3), (3, 0), (3, 1), (3, 2), (3, 3)])
We can now instantiate our PoolingWindows object (with the same number of scales as the pyramid) and pass it our new dictionary. The pooled_coeffs will have the same keys as our input dictionary and its values will be pooled versions of the corresponding coefficients with the 3rd dimension corresponding to each window. Note the warnings about some windows being too small!
pw = fen.PoolingWindows(0.5, img.shape[-2:], num_scales=pyr.num_scales)
pooled_coeffs = pw(new_pyr)
for k, v in pooled_coeffs.items():
print(f'scale {k[0]}, orientation band {k[1]}: {v.shape}')
scale 0, orientation band 0: torch.Size([1, 1, 1380])
scale 0, orientation band 1: torch.Size([1, 1, 1380])
scale 0, orientation band 2: torch.Size([1, 1, 1380])
scale 0, orientation band 3: torch.Size([1, 1, 1380])
scale 1, orientation band 0: torch.Size([1, 1, 1380])
scale 1, orientation band 1: torch.Size([1, 1, 1380])
scale 1, orientation band 2: torch.Size([1, 1, 1380])
scale 1, orientation band 3: torch.Size([1, 1, 1380])
scale 2, orientation band 0: torch.Size([1, 1, 1380])
scale 2, orientation band 1: torch.Size([1, 1, 1380])
scale 2, orientation band 2: torch.Size([1, 1, 1380])
scale 2, orientation band 3: torch.Size([1, 1, 1380])
scale 3, orientation band 0: torch.Size([1, 1, 1380])
scale 3, orientation band 1: torch.Size([1, 1, 1380])
scale 3, orientation band 2: torch.Size([1, 1, 1380])
scale 3, orientation band 3: torch.Size([1, 1, 1380])
/home/docs/checkouts/readthedocs.org/user_builds/pooling-windows/envs/latest/lib/python3.12/site-packages/fenestration/pooling_windows.py:284: UserWarning: Creating windows for scale 1 with min_ecc 0.5, but calculated minimal eccentricity is 0.7480167757526863, so be aware some are smaller than a pixel!
warnings.warn(
/home/docs/checkouts/readthedocs.org/user_builds/pooling-windows/envs/latest/lib/python3.12/site-packages/fenestration/pooling_windows.py:284: UserWarning: Creating windows for scale 2 with min_ecc 0.5, but calculated minimal eccentricity is 1.4960335515053724, so be aware some are smaller than a pixel!
warnings.warn(
/home/docs/checkouts/readthedocs.org/user_builds/pooling-windows/envs/latest/lib/python3.12/site-packages/fenestration/pooling_windows.py:284: UserWarning: Creating windows for scale 3 with min_ecc 0.5, but calculated minimal eccentricity is 2.9920671030107453, so be aware some are smaller than a pixel!
warnings.warn(