Run this notebook yourself!

Download the executed notebook: deep_nets.ipynb!

Run it in your browser: deep_nets.ipynb!

Using Deep Neural Networks with plenoptic#

Warning

This notebook requires the optional dependency torchvision, which can be installed with pip.

Plenoptic is compatible with any model written in pytorch, including deep neural networks. In this notebook we show how to use plenoptic with models from the deep network zoos TorchVision and timm, creating a metamer for an intermediate layer of ResNet50.

You may also be interested in Reproducing ResNet50 Image Metamers from Feather et al., 2023, where we create model metamers for several ResNet50 intermediate layers, reproducing some of the results from Feather et al., 2023.

import matplotlib.pyplot as plt
import numpy as np
import torch

import plenoptic as po

# this notebook uses torchvision, which is an optional dependency.
# if this import fails, install torchvision in your plenoptic
# environment and restart the notebook kernel.
try:
    import torchvision
except ModuleNotFoundError:
    raise ModuleNotFoundError(
        "optional dependency torchvision not found!"
        " please install it in your plenoptic environment "
        "and restart the notebook kernel"
    )


DEVICE = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# so that relative sizes of axes created by po.plot.imshow and others look right
plt.rcParams["figure.dpi"] = 72
# Animation-related settings
plt.rcParams["animation.html"] = "html5"
# use single-threaded ffmpeg for animation writer
plt.rcParams["animation.writer"] = "ffmpeg"
plt.rcParams["animation.ffmpeg_args"] = ["-threads", "1"]

# set seed for reproducibility
po.set_seed(1)

Prepare model and image for synthesis#

In this section, we walk through how to initialize a plenoptic-compatible model using the weights from TorchVision. Then, at the end of this section, we briefly show to do the same with models from timm.

To use one of these deep nets in plenoptic, we have to specify three things:

  1. The deep net model.

  2. The layer(s) to extract.

  3. The image pre-processing to use.

Initialize deep neural network and pre-trained weights#

First, we download the model weights for ResNet50 trained on ImageNet-1K and initialize the torchvision model.

weights = torchvision.models.ResNet50_Weights.IMAGENET1K_V1
deepnet = torchvision.models.resnet50(weights=weights)

Next, we ensure that our model is in evaluation mode. Many models, including ResNet50, behave differently when in training and evaluation mode. In plenoptic, models are fixed and so we want the evaluation behavior (see here for more details):

deepnet.eval();

Select layer#

Next, we specify the layer to target. Figure 2e from Feather et al., 2023 shows an interesting progression for ResNet50 metamers: the layer 2 metamer looks almost like the target image, the layer 3 metamer shows some RGB noise, and the layer 4 metamer is almost completely unidentifiable.

Let’s start with "layer3". See Reproducing ResNet50 Image Metamers from Feather et al., 2023 for the other layers, and note that you can specify multiple layers simultaneously.

target_layer = "layer3"

Specify preprocessing#

Finally, it is important to specify the preprocessing transform of the model. During neural network training, images are transformed before being passed to the network, and the same transformation needs to be applied when using the trained network. As the torchvision docs explain it (quoting version 0.27):

Before using the pre-trained models, one must preprocess the image (resize with right resolution/interpolation, apply inference transforms, rescale the values etc). There is no standard way to do this as it depends on how a given model was trained. It can vary across model families, variants or even weight versions. Using the correct preprocessing method is critical and failing to do so may lead to decreased accuracy or incorrect outputs.

For models trained on ImageNet, this preprocessing consists of two steps: resizing to a height and width of 224 pixels and normalizing the color channels (subtracting means and dividing by standard deviations). Following Feather et al., 2023, we recommend including the normalization step in the model for metamer synthesis, but handling the image resizing externally.

In torchvision, this preprocessing transform is a single torch.nn.Module which we cannot easily subdivide. This module includes resizing, cropping and normalization:

transform = weights.transforms()
print(transform)
ImageClassification(
    crop_size=[224]
    resize_size=[256]
    mean=[0.485, 0.456, 0.406]
    std=[0.229, 0.224, 0.225]
    interpolation=InterpolationMode.BILINEAR
)

Since we cannot grab a subsection of this transform, we instead create a separate normalization transform, using the specified mean and std:

norm = torchvision.transforms.Normalize(transform.mean, transform.std)

Prepare the image#

Now, let’s prepare the image. The input image needs to be an RGB image with a height and width of 224 pixels. It should probably also be like those found in ImageNet: a single object in the center of the frame that belongs to one of the image classes. We’ll use one of the famous monkey selfies, and resize it appropriately:

img = po.data.macaque()
# here we downsample the original image by a factor of 4 and then lop off the bottom.
# that way, when we take the central 224 pixels in the following block, we end up with a
# decent image.
img = po.process.blur_downsample(img, 2)[..., :-59, :]

As discussed above, models trained on ImageNet should be passed an image of size 224 by 224. We’ll use plenoptic’s plenoptic.process.center_crop to do so, grabbing the required size directly from the model’s associated transform;

img = po.process.center_crop(img, transform.crop_size[0])
po.plot.imshow(img, as_rgb=True);
../../_images/e8788eb276e53e63cff9a1fb1250dfe75cd6dc8f27aadc370a01b930e7a34d4c.png

Last steps#

Now we can finally create our model by passing the neural network, target layer, and preprocessing transform to plenoptic’s DeepNetFeatures

model = po.models.DeepNetFeatures(deepnet, target_layer, norm)

Finally, let’s remove the gradient from all model parameters (as models in plenoptic are fixed), convert everything to float64, for reproducibility, and move everything to DEVICE:

img = img.to(DEVICE).to(torch.float64)
model.to(DEVICE).to(torch.float64)
po.remove_grad(model)

Understand the model#

Image classification#

ResNet50 is trained to classify images into one of 1000 categories. Any metamer of an intermediate layer should preserve this classification, which is the output of the final layer; this is one of the criteria that Feather et al., 2023 check for synthesis success. Let’s examine that classification now, creating a little helper function:

imagenet_categories = np.asarray(weights.meta["categories"])
# Move deepnet to float64, device since we know our images are all float64
deepnet.to(torch.float64).to(DEVICE)


def get_category(image):
    # Get probabilities of each image category
    image_cat = torch.nn.functional.softmax(deepnet(norm(image)), dim=1)
    # Convert to 1d numpy array, so we can use as index in
    # imagenet_categories above.
    image_cat = po.to_numpy(image_cat.squeeze())
    return imagenet_categories[image_cat.argmax()]


print(f"ResNet50 predicted class: {get_category(img)}")
ResNet50 predicted class: guenon

The category, guenon, is an Old World monkey. Though it isn’t the actual species of the monkey in question (a Celebes crested macaque), it’s a reasonable category for it.

After we synthesize the metamer, we will ensure that our model correctly classifies it as a guenon as well.

Output visualization#

Our model object now returns only the activations from our specified layer as a single 2d vector (with the first dimension corresponding to the batch dimension of our input):

rep = model(img)
print(rep)
print(rep.shape)
tensor([[0.1338, 0.0000, 0.0897,  ..., 0.3310, 0.0000, 0.0000]],
       device='cuda:0', dtype=torch.float64)
torch.Size([1, 200704])

We have flattened the model representation of the given layer (to support representations from multiple layers simultaneously). If you would like to retrieve the original shape, you can use the convert_to_dict method:

rep = model.convert_to_dict(rep)
print(rep.keys())
print(rep[target_layer].shape)
odict_keys(['layer3'])
torch.Size([1, 1024, 14, 14])

DeepNetFeatures also has a plot_representation method, which creates two subplots. The first plots the average across channel, the average spatial representation, while the second averages across space to get a per-channel average representation:

fig, _ = model.plot_representation(rep)
../../_images/12cf0106cad3189f8c1185c0c1dccb9d5d83955642dbc99f064468b0a717d117.png

Now that we understand our model, let’s synthesize some metamers!

Synthesize the metamer#

Let us initialize our metamer object using the above image and model:

met = po.Metamer(img, model)

We could just call synthesize to synthesize the metamer from here, but we have found better results when tweaking the optimization hyperparameters. We can do this using setup:

Warning

The hyperparameters shown here are the ones that work best for this deep net, target layer, and image. You will likely need to tweak them for other synthesis problems. It is likely that they make a reasonable starting point, but you are encouraged to experiment!

scheduler = torch.optim.lr_scheduler.StepLR
scheduler_kwargs = {"step_size": 3000, "gamma": 0.5}
met.setup(
    optimizer_kwargs={"amsgrad": False},
    scheduler=scheduler,
    scheduler_kwargs=scheduler_kwargs,
)

In the above:

  • We are using one of torch’s learning-rate schedulers, StepLR, to halve (gamma) the learning-rate every 3000 steps (step_size).

  • We are not changing the optimizer or its learning rate from their default values (Adam and 0.01, respectively, see setup for details on other defaults).

  • We are, however, specifying that we do not want to use the AMSGrad variant of the Adam optimizer, which is plenoptic’s default behavior.

  • We are also not changing our initial image from the default: random pixels uniformly distributed between 0 and 1.

Now that we’ve set our optimization hyperparameters, we can synthesize our metamer!

# by setting stop_iters_to_check=max_iter, we ensure it keeps going through
# all iterations
met.synthesize(max_iter=6000, stop_iters_to_check=6000, store_progress=100)

How many iterations?

Here we’re only running optimization for 6000 iterations, which is enough to demonstrate the point, but if you were to use these metamers in an experiment, we would recommend running synthesis for longer and thinking carefully about your success criteria, see What is a “good” metamer? and Synthesis success? for more discussion.

Let’s call synthesis_status to visualize the synthesis status:

po.plot.synthesis_status(met, figsize=(15, 4.5));
Clipping input data to the valid range for imshow with RGB data ([0..1] for floats or [0..255] for integers). Got range [-0.0016976601424805965..1.0000506869582761].
../../_images/80789da1d586af61ac38c2657ce1ae6e65fc26c9f9146353c77943dde1a8de5e.png

In the above plot, we can see the metamer in the leftmost subplot, the loss over synthesis iterations in the middle, and the representation error on the right:

  • Our metamer matches the result discussed earlier in this notebook: as a layer 3 metamer, it looks similar to the original image with some RGB noise added.

  • We can see that the optimization performed reasonably well: the loss decreased gradually over synthesis. If you were using these stimuli in an experiment, it would be worth continuing a bit more to get the loss even lower, but this demonstrates the point.

  • The representation error plot has the same structure as the plot_representation plot above. We can see that, while there’s variation across both channels and space, there’s not an obvious outlier whose error we have been unable to constrain.

And we can animate the above figure over synthesis iterations as well, to see the metamer take shape:

po.plot.synthesis_animate(met, figsize=(15, 4.5))
/home/jenkins/agent/workspace/CCN_neurorse_plenoptic_PR-460/lib/python3.12/site-packages/plenoptic/plot/synthesis.py:2074: UserWarning: 60 frames of saved metamer clipped: Clipping input data to the valid range for imshow with RGB data ([0..1] for floats). To avoid clipping, use process_image argument.
  warnings.warn(warning_msg)

Finally, let’s ensure that our metamer has the same image category as our initial image:

print(f"Target image category: {get_category(met.image)}")
print(f"Metamer category: {get_category(met.metamer)}")
Target image category: guenon
Metamer category: guenon

In this notebook, we have demonstrated how to use deep neural networks from external models zoos with plenoptic.models.DeepNetFeatures, and shown how to generate metamers for an intermediate layer.