# reconstruction_segmentation_loss

Python API: `luxonis_train.attached_modules.losses.reconstruction_segmentation_loss`

The DRAEM loss, which adds the L2 and the structural similarity errors of the reconstruction to a focal loss over the anomaly
mask.

The module also holds the structural similarity (SSIM) helpers that the loss uses.

## Classes

### ReconstructionSegmentationLoss

Reconstruction and anomaly segmentation loss for
[DiscSubNetHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/discsubnet_head/discsubnet_head.md).

The loss follows [DRAEM](https://arxiv.org/abs/2108.07610). It compares the reconstructed image with the clean image, and the
predicted anomaly segmentation with the anomaly mask.

 * `Inputs:`: * `predictions` (`Tensor`): [B, 2, H, W] anomaly logits
    * `reconstruction` (`Tensor`): [B, 3, H, W] reconstructed image
    * `target_original_segmentation` (`Tensor`): [B, 3, H, W] clean image
    * `target_segmentation` (`Tensor`): [B, 2, H, W] one-hot anomaly mask
 * `Outputs:`: * `Tensor`: scalar total loss, or [B, H, W] when `reduction` is `"none"`
    * `dict[str, Tensor]`: sub-losses `l2_loss`, `ssim_loss`, `focal_loss`
 * `Formula:`: r is the reconstruction, x the clean image, z the logits, and y the one-hot mask: L = MSE(r, x) + (1 − SSIM(r, x))
   + FL(z, y) MSE is the mean over all elements. SSIM is the mean of the map from
   [ssim](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/reconstruction_segmentation_loss.md),
   with an 11×11 Gaussian window. FL is
   [SoftmaxFocalLoss](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/softmax_focal_loss.md)
   with `alpha`, `gamma`, `smooth`, and `reduction`.

> **References**
> * Source: This project.
 * License: Apache-2.0 (this project)

> **Notes**
> The focal term comes from an internal [SoftmaxFocalLoss](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/softmax_focal_loss.md) without a node. The SSIM term runs in `float32` with autocast off, and it estimates the dynamic range from `reconstruction`. The sub-losses are not detached.

> **Example**
> Attached to a `DiscSubNetHead` in `model.nodes`:

```yaml
- name: DiscSubNetHead
  inputs: [RecSubNet]
  losses:
    - name: ReconstructionSegmentationLoss
```

 * `Compatible with:`: * Used by:
   [AnomalyDetectionModel](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/config/predefined_models/anomaly_detection/v1/model.md)
    * Nodes:
      [DiscSubNetHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/discsubnet_head/discsubnet_head.md)

#### Methods

##### init

```python
def __init__(alpha: float = 1, gamma: float = 2.0, reduction: Literal['none', 'mean', 'sum'] = 'mean', smooth: float = 1e-05, **kwargs):
```

Initialize the L2, SSIM, and focal losses.

Parameters

 * `alpha` (`float`): The factor α of
   [SoftmaxFocalLoss](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/softmax_focal_loss.md).
   It scales the focal loss of every pixel.
 * `gamma` (`float`): The focal exponent γ of
   [SoftmaxFocalLoss](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/softmax_focal_loss.md).
   A larger value gives less weight to the pixels that the head already classifies well.
 * `reduction` (`Literal['none', 'mean', 'sum']`): How
   [SoftmaxFocalLoss](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/softmax_focal_loss.md)
   reduces the loss of the pixels. With `"none"`, the focal term and the total loss have the shape `[B, H, W]`.
 * `smooth` (`float`): The label smoothing of
   [SoftmaxFocalLoss](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/softmax_focal_loss.md).
   Its constructor raises `ValueError` when `smooth` is not in `[0, 1]`.
 * `**kwargs`: Keyword arguments forwarded to
   [BaseLoss](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/base_loss.md),
   such as `node` and `final_loss_weight`.

##### forward

```python
def forward(predictions: Tensor, reconstruction: Tensor, target_original_segmentation: Tensor, target_segmentation: Tensor) -> tuple[Tensor, dict[str, Tensor]]:
```

Compute the reconstruction and segmentation losses.

[SoftmaxFocalLoss](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/softmax_focal_loss.md)
raises `ValueError` when `predictions` has fewer than `2` channels, or when its shape is not the shape of `target_segmentation`.

> **Example**
> A perfect reconstruction makes the L2 and SSIM terms `0`. Zero logits give the focal term − (1 − 0.5)2ln0.5 ≈ 0.1733:

```pycon
>>> import torch
>>> loss = ReconstructionSegmentationLoss()
>>> image = torch.linspace(0, 1, 48).reshape(1, 3, 4, 4)
>>> logits = torch.zeros(1, 2, 4, 4)
>>> mask = torch.zeros(1, 2, 4, 4)
>>> mask[:, 0] = 1
>>> total, sub_losses = loss(logits, image, image, mask)
>>> sorted(sub_losses)
['focal_loss', 'l2_loss', 'ssim_loss']
>>> round(total.item(), 4)
0.1733
```

Parameters

 * `predictions` (`Tensor`): Anomaly logits of shape `[B, C, H, W]`, with `C` at least `2`. The `segmentation` output of the node.
 * `reconstruction` (`Tensor`): Reconstructed images of shape `[B, 3, H, W]`. The `reconstruction` output of the node.
 * `target_original_segmentation` (`Tensor`): Clean images of the shape of `reconstruction`. The `original_segmentation` label of
   the task.
 * `target_segmentation` (`Tensor`): One-hot anomaly masks of the shape of `predictions`. The `segmentation` label of the task.

Returns

 * `tuple[Tensor, dict[str, Tensor]]`: The sum of the three terms, and a dictionary that maps `"l2_loss"`, `"ssim_loss"`, and
   `"focal_loss"` to the terms. The L2 and SSIM terms are scalars. The focal term and the sum are scalars, or of shape `[B, H, W]`
   when `reduction` is `"none"`.

#### Attributes

##### loss_focal

##### loss_l2

##### loss_ssim

##### node

##### supported_tasks

### SSIM

Structural dissimilarity loss, 1 − SSIM.

The module calls
[ssim](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/reconstruction_segmentation_loss.md)
with a Gaussian window of the standard deviation `1.5` and returns one minus the result. It caches the window for the channel
count of the last input.

#### Methods

##### init

```python
def __init__(window_size: int = 11, size_average: bool = True, val_range: float | None = None):
```

Initialize the loss with a window for one channel.

Parameters

 * `window_size` (`int`): Side of the square Gaussian window, in pixels.
 * `size_average` (`bool`): Whether
   [forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/reconstruction_segmentation_loss.md)
   averages over the whole batch. When `False`, it returns one value for each image.
 * `val_range` (`float | None`): A dynamic range for
   [ssim](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/reconstruction_segmentation_loss.md).
   [forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/reconstruction_segmentation_loss.md)
   does not pass it to
   [ssim](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/reconstruction_segmentation_loss.md),
   so
   [ssim](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/reconstruction_segmentation_loss.md)
   always estimates the dynamic range from `img1`.

##### forward

```python
def forward(img1: Tensor, img2: Tensor) -> Tensor:
```

Return one minus the SSIM of two image batches.

The method turns autocast off and casts both batches to `float32`. When `img1` has another channel count than the cached window,
the method builds a new window and caches it.

> **Example**
> ```pycon
>>> import torch
>>> loss = SSIM()
>>> image = torch.linspace(0, 1, 768).reshape(1, 3, 16, 16)
>>> loss(image, image).item()
0.0
>>> round(loss(image, 1 - image).item(), 4)
0.7642
```

Parameters

 * `img1` (`Tensor`): Images of shape `[B, C, H, W]`. [ssim](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/reconstruction_segmentation_loss.md) estimates the dynamic range from them.
 * `img2` (`Tensor`): Images of the shape of `img1`.

Returns

 * `Tensor`: 1 − SSIM as a scalar, or of shape `[B]` when `size_average` is `False`. `0` for equal images.

## Functions

### create_window

```python
def create_window(window_size: int, channel: int = 1) -> Tensor:
```

Build a normalized 2D Gaussian window for [ssim](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/reconstruction_segmentation_loss.md).

The window is the outer product of a 1D [gaussian](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/reconstruction_segmentation_loss.md) with itself. The standard deviation is always `1.5`. Each channel gets the same window.

> **Example**
> ```pycon
>>> import torch
>>> window = create_window(11, channel=3)
>>> window.shape
torch.Size([3, 1, 11, 11])
>>> torch.allclose(window.sum(dim=(1, 2, 3)), torch.ones(3))
True
```

Parameters

 * `window_size` (`int`): Side of the square window, in pixels.
 * `channel` (`int`): Number of channels.

Returns

 * `Tensor`: A `float32` window of shape `[channel, 1, window_size, window_size]`, the weight shape of a grouped convolution. The
   weights of each channel sum to `1`.

### gaussian

```python
def gaussian(window_size: int, sigma: float) -> Tensor:
```

Build a normalized 1D Gaussian of `window_size` samples.

Sample i gets the weight exp⎛⎝ − ((i − c)2)/(2σ2)⎞⎠, with the center c at `window_size // 2`. The function then divides the
weights by their sum. For an even `window_size`, the center is the right one of the two middle samples.

> **Example**
> ```pycon
>>> [round(weight, 4) for weight in gaussian(3, 1.5).tolist()]
[0.3078, 0.3844, 0.3078]
```

Parameters

 * `window_size` (`int`): Number of samples.
 * `sigma` (`float`): Standard deviation σ, in samples.

Returns

 * `Tensor`: The weights, of shape `[window_size]`, with the sum `1`.

### ssim

```python
def ssim(img1: Tensor, img2: Tensor, window_size: int = 11, window: Tensor | None = None, size_average: bool = True, val_range:
float | None = None) -> Tensor:
```

Compute the structural similarity (SSIM) of two image batches.

The window filters each channel of both batches, with zero padding. The filter gives the local means μ1 and μ2, the local variances σ21 and σ22, and the local covariance σ12. The SSIM map is

SSIM = ((2μ1μ2 + C1)(2σ12 + C2))/((μ21 + μ22 + C1)(σ21 + σ22 + C2))

with C1 = (0.01L)2 and C2 = (0.03L)2 for the dynamic range L.

> **Example**
> ```pycon
>>> import torch
>>> image = torch.linspace(0, 1, 64).reshape(1, 1, 8, 8)
>>> ssim(image, image).item()
1.0
>>> round(ssim(image, 1 - image).item(), 4)
0.0926
```

Parameters

 * `img1` (`Tensor`): Images of shape `[B, C, H, W]`.
 * `img2` (`Tensor`): Images of the shape of `img1`.
 * `window_size` (`int`): Side of the window. The padding is always `window_size // 2`, also when `window` has another size.
 * `window` (`Tensor | None`): Filter weights of shape `[C, 1, k, k]`, for example from
   [create_window](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/reconstruction_segmentation_loss.md).
   When `None`, the function builds a window with the side `min(window_size, H, W)`.
 * `size_average` (`bool`): Whether to return the mean of the whole map. When `False`, the function returns the mean of each
   image.
 * `val_range` (`float | None`): The dynamic range L. When `None`, the function estimates L from `img1` as the maximum minus the
   minimum. The maximum is `255` when a value of `img1` is above `128`, otherwise `1`. The minimum is `-1` when a value of `img1`
   is below `-0.5`, otherwise `0`.

Returns

 * `Tensor`: The mean SSIM as a scalar, or of shape `[B]` when `size_average` is `False`. `1` for equal images.
