# gradcam_visualizer

Python API: `luxonis_train.callbacks.gradcam_visualizer`

Logs Grad-CAM heat maps of validation images.

A heat map shows the image regions that contribute to the score of one class.

## Classes

### GradCamCallback

Callback that logs Grad-CAM heat maps of validation images.

The callback is experimental. It runs `HiResCAM` of `pytorch_grad_cam` on the first `log_n_batches` batches of each validation
epoch. It puts each heat map over its image and logs the result with the tracker of the model. It logs nothing outside of
validation.

Grad-CAM needs gradients. The validation loop of a fit runs without inference mode, so the callback works there.
`trainer.validate` runs in inference mode by default. There, the backward pass of Grad-CAM raises `RuntimeError`.

The callback is in the `CALLBACKS` registry, so a config can add it:

```yaml
trainer:
  callbacks:
    - name: GradCamCallback
      params:
        target_layer: 10
        task: segmentation
```

#### Methods

##### init

```python
def __init__(target_layer: int, class_idx: int = 0, log_n_batches: int = 1, task: str = 'classification'):
```

Initialize the callback.

Parameters

 * `target_layer` (`int`): The index of the layer that Grad-CAM reads, in the order of `named_modules()` of the
   [PLModuleWrapper](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/gradcam_visualizer.md).
   Index `0` is the wrapper and index `1` is the wrapped model. The callback selects the layer with the slice `[target_layer :
   target_layer + 1]`. When the slice is empty, for example for `-1` or an index out of range, Grad-CAM raises `ValueError`.
 * `class_idx` (`int`): The index of the class that the heat maps explain.
 * `log_n_batches` (`int`): The number of batches to log in each validation epoch, from the first batch.
 * `task` (`str`): The type of the output to explain. One of `"segmentation"`, `"detection"`, `"classification"`, or
   `"keypoints"`. See
   [PLModuleWrapper.forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/gradcam_visualizer.md).

##### on_validation_batch_end

```python
def on_validation_batch_end(trainer: pl.Trainer, pl_module: lxt.LuxonisLightningModule, outputs: STEP_OUTPUT, batch: tuple[dict[str, Tensor], Packet[Tensor]], batch_idx: int):
```

Log the heat maps of the first `log_n_batches` batches.

Lightning calls this hook after every validation batch. When `batch_idx` is lower than `log_n_batches`, the hook takes the inputs
of the batch. From a dictionary of inputs, it takes the entry `pl_module.image_source`. Then it calls
[GradCamCallback.visualize_gradients](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/gradcam_visualizer.md).
For a later batch, the hook does nothing.

Parameters

 * `trainer` (`pl.Trainer`): The trainer. It gives the step of the logged images.
 * `pl_module` (`lxt.LuxonisLightningModule`): The model. It gives `image_source`, the config, and the tracker.
 * `outputs` (`STEP_OUTPUT`): The output of the validation step. Unused.
 * `batch` (`tuple[dict[str, Tensor], Packet[Tensor]]`): The inputs and the labels of the batch.
 * `batch_idx` (`int`): The index of the batch in the validation epoch.

##### setup

```python
def setup(trainer: pl.Trainer, pl_module: lxt.LuxonisLightningModule, stage: str):
```

Wrap the model for Grad-CAM.

Lightning calls this hook at the start of every stage. The hook creates a new
[PLModuleWrapper](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/gradcam_visualizer.md)
of `pl_module` and `task` for
[visualize_gradients](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/gradcam_visualizer.md).

Parameters

 * `trainer` (`pl.Trainer`): The trainer. Unused.
 * `pl_module` (`lxt.LuxonisLightningModule`): The model to wrap.
 * `stage` (`str`): The stage that starts. Unused.

##### visualize_gradients

```python
def visualize_gradients(trainer: pl.Trainer, pl_module: lxt.LuxonisLightningModule, images: Tensor, batch_idx: int):
```

Compute the Grad-CAM heat maps of a batch and log them.

The method creates a `HiResCAM` on the layer at index `target_layer` in `named_modules()` of the
[PLModuleWrapper](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/gradcam_visualizer.md).
`HiResCAM` sets the wrapper and the model to eval mode. The method keeps the `HiResCAM`, so its hooks stay on the layer until a
later call replaces it. Until then, the hooks keep a CPU copy of the output of the layer from each forward pass, and of its
gradient from each backward pass.

The Grad-CAM target of each image depends on `task`:

 * `"segmentation"`: the method first runs the model on the images and applies a softmax over the classes. The mask of an image
   holds the pixels where `class_idx` has the highest probability. The target is the sum of the `class_idx` score map over that
   mask.
 * Any other task: the target is the `class_idx` entry of the scores from
   [PLModuleWrapper.forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/gradcam_visualizer.md).

The method computes the heat maps with gradients enabled. Grad-CAM clears the gradients of the model and runs a backward pass, so
the parameters of the model keep new `grad` values. The method denormalizes the images with `trainer.preprocessing.normalize` of
`pl_module.cfg`. Each heat map goes over its image as a JET color map at half opacity. The tracker of `pl_module` logs each result
as `gradcam/gradcam_<batch_idx>_<i>` at `trainer.global_step`. `<i>` is the index of the image in the batch.

Parameters

 * `trainer` (`pl.Trainer`): The trainer. It gives the step of the logged images.
 * `pl_module` (`lxt.LuxonisLightningModule`): The model. It gives the config for the denormalization and the tracker.
 * `images` (`Tensor`): The normalized images, of shape `[B, C, H, W]`.
 * `batch_idx` (`int`): The index of the batch. It is part of the image names.

### PLModuleWrapper

Lightning module that gives Grad-CAM one score tensor per batch.

Grad-CAM needs a model that takes a tensor of images and returns a tensor of scores.
[LuxonisLightningModule.full_forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
returns a
[LuxonisOutput](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_output.md)
with the packets of the output nodes. The wrapper takes the images and returns one tensor from these packets.
[GradCamCallback](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/gradcam_visualizer.md)
creates the wrapper.

#### Methods

##### init

```python
def __init__(pl_module: lxt.LuxonisLightningModule, task: str):
```

Initialize the wrapper.

Parameters

 * `pl_module` (`lxt.LuxonisLightningModule`): The model to wrap.
 * `task` (`str`): Selects the output that
   [PLModuleWrapper.forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/gradcam_visualizer.md)
   returns. One of `"segmentation"`, `"detection"`, `"classification"`, or `"keypoints"`. The constructor does not check the
   value.

##### forward

```python
def forward(inputs: Tensor, *args, **kwargs) -> Tensor:
```

Run the model on images and return the scores for `task`.

The method calls
[LuxonisLightningModule.full_forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
with `{"image": inputs}` and the extra arguments. When the model has more than one output node, the method logs a warning. It
reads the packet of the first output node. The returned tensor depends on `task`:

 * `"segmentation"`: the `"segmentation"` output, unchanged.
 * `"classification"`: the `"classification"` output, unchanged.
 * `"detection"` and `"keypoints"`: the `"class_scores"` output, summed over dimension 1. For the scores of
   [EfficientBBoxHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/efficient_bbox_head.md),
   of shape `[B, N, n_classes]` for `N` anchors, the result has the shape `[B, n_classes]`.

Parameters

 * `inputs` (`Tensor`): The images, of shape `[B, C, H, W]`. The method passes them as the input named `"image"`. The name does
   not follow `loader.image_source` of the config.
 * `*args`: Extra positional arguments for
   [LuxonisLightningModule.full_forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md).
 * `**kwargs`: Extra keyword arguments for
   [LuxonisLightningModule.full_forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md).

Returns

 * `Tensor`: The scores for `task`.

Raises

 * `ValueError`: When `task` is not one of the four supported values.
