# fomo_localization_loss

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

The weighted focal loss of the FOMO heatmap of object centers.

## Classes

### FOMOLocalizationLoss

Focal loss of
[FOMOHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/fomo_head.md)
over a heatmap of object centers.

 * `Inputs:`: * `heatmap` (`Tensor`): [B, nclasses, Hf, Wf] logits
    * `target` (`Tensor`): [M, 6], `[batch, class, x, y, w, h]`, `xywh` normalized
 * `Outputs:`: * `Tensor`: scalar
 * `Formula:`: The target heatmap y is `1` in the cell that holds the center of a box, in the channel of the box class. It is `0`
   in all other cells. For the logit x of each cell: L = (1)/(B nclasses Hf Wf)∑α (1 − pt)γ w BCE(x, y), pt = exp( − BCE(x, y))
   BCE is the binary cross entropy with logits, and pt is the predicted probability of the target value. w is `object_weight` in a
   center cell and `1` in the other cells. α is `alpha`, and γ is `gamma`.

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

> **Notes**
> The FOMO task has no keypoint labels, so the loss uses the box centers as the targets. The center of a box with the normalized top-left corner (x, y) and size (w, h) is in column ⌊(x + w ⁄ 2) Wf⌋ and row ⌊(y + h ⁄ 2) Hf⌋. Two boxes of one class with centers in the same cell give one target cell. Unlike the usual focal loss, α scales the center cells and the other cells by the same factor.

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

```yaml
- name: FOMOHead
  inputs: [EfficientRep]
  losses:
    - name: FOMOLocalizationLoss
```

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

#### Methods

##### init

```python
def __init__(object_weight: float = 500, alpha: float = 0.45, gamma: float = 2, **kwargs):
```

Initialize the loss.

The method reads the input image size of the node.
[forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/fomo_localization_loss.md)
does not use the size. The loss therefore needs a node: without `node`,
[BaseAttachedModule.node](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/base_attached_module.md)
raises `RuntimeError`.

Parameters

 * `object_weight` (`float`): The factor w of the loss in a center cell. The other cells have the factor `1`.
 * `alpha` (`float`): The factor α of the loss in every cell.
 * `gamma` (`float`): The exponent γ of the focal term. A higher value lowers the loss of the cells that the head already predicts
   well. `0` removes the focal term.
 * `**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 `final_loss_weight` and `node`.

##### forward

```python
def forward(heatmap: Tensor, target: Tensor) -> Tensor:
```

Build the target heatmap and compute the weighted focal loss.

> **Example**
> The example uses a heatmap of zero logits on a `4x4` grid. One box has its center in column `2` and row `2`. Each cell has the binary cross entropy ln2 and the focal factor 0.45⋅0.52. The center cell has the weight `500`, so it gives most of the loss:

```pycon
>>> import torch
>>> from torch import Size
>>> from luxonis_train.nodes import FOMOHead
>>> sizes = [Size([1, 8, 8, 8]), Size([1, 16, 4, 4])]
>>> head = FOMOHead(
...     n_classes=1,
...     input_shapes=[{"features": sizes}],
...     original_in_shape=Size([3, 32, 32]),
... )
>>> heatmap = torch.zeros(1, 1, 4, 4)
>>> target = torch.tensor([[0.0, 0.0, 0.25, 0.25, 0.5, 0.5]])
>>> loss = FOMOLocalizationLoss(node=head)
>>> round(loss(heatmap, target).item(), 4)
2.51
>>> loss = FOMOLocalizationLoss(node=head, object_weight=1)
>>> round(loss(heatmap, target).item(), 4)
0.078
```

Parameters

 * `heatmap` (`Tensor`): The class logits of shape `[B, n_classes, H, W]`, from the `heatmap` key of the node output.
 * `target` (`Tensor`): The `boundingbox` label of shape `[M, 6]`. Each row holds the batch index, the class, and the normalized
   `x`, `y`, `w`, and `h` of one box. `x` and `y` give the top-left corner.

Returns

 * `Tensor`: The mean of the weighted focal loss over all cells of `heatmap`, as a scalar.

#### Attributes

##### node

##### supported_tasks
