# ctc_loss

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

The connectionist temporal classification (CTC) loss, which trains a sequence model without an alignment between the input and the
text.

## Classes

### CTCLoss

CTC loss for OCR, with an optional focal factor.

 * `Inputs:`: * `predictions` (`Tensor`): [B, T, C] logits, class `0` is blank
    * `target` (`Tensor`): [B, L] Unicode code points of the text, `0`-padded
 * `Outputs:`: * `Tensor`: scalar
 * `Formula:`: The encoder of the node maps each character of `target` to its class index. For image i, ℓi is the negative
   log-likelihood of its encoded text. `nn.CTCLoss` sums the likelihood over all alignments of the T log-softmax predictions, with
   class `0` as blank. With `use_focal_loss`, the loss is L = (1)/(B)B∑i = 1ℓi(1 − e − ℓi)2 Without `use_focal_loss`, L is the
   mean of the ℓi.

> **References**
> * Source: Wraps [torch.nn.CTCLoss](https://docs.pytorch.org/docs/stable/generated/torch.nn.CTCLoss.html) (BSD-3-Clause).
 * License: Apache-2.0 (this project)

> **Notes**
> The loss declares no `supported_tasks` and takes its task from the node. Its `node` annotation limits it to [OCRCTCHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/ocr_ctc_head.md), and [BaseAttachedModule](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/base_attached_module.md) raises [IncompatibleError](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/utils/exceptions.md) for another node. The loss does not detach the focal factor, so the gradient also flows through it. `nn.CTCLoss` keeps `zero_infinity` off, so a text that no alignment of length T can produce gives an infinite loss.

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

```yaml
- name: OCRCTCHead
  inputs: [SVTRNeck]
  losses:
    - name: CTCLoss
```

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

#### Methods

##### init

```python
def __init__(use_focal_loss: bool = True, **kwargs):
```

Initialize the loss and the wrapped `nn.CTCLoss`.

The wrapped loss uses class `0` as blank and returns the loss of each image.
[forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/ctc_loss.md)
takes the mean over the batch.

Parameters

 * `use_focal_loss` (`bool`): Whether to multiply the loss ℓi of each image by (1 − e − ℓi)2, as the class formula describes.
 * `**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(predictions: Tensor, target: Tensor) -> Tensor:
```

Compute the CTC loss of a batch of text predictions.

The method encodes `target` with the encoder of the node and moves the result to the device of `predictions`. The encoder drops a
character that is not in the alphabet. When `ignore_unknown` of the node is `False`, it maps the character to `<UNK>` instead. The
length of each text is its number of non-zero class indices. Every prediction sequence has the full length `T`.

> **Example**
> The example has one time step with equal logits for the blank and the classes of `"a"` and `"b"`. The text `"a"` has one alignment with the probability 1 ⁄ 3, so ℓ = ln3. The focal factor is (1 − 1 ⁄ 3)2 = 4 ⁄ 9:

```pycon
>>> import torch
>>> from luxonis_train.nodes import OCRCTCHead
>>> head = OCRCTCHead(
...     alphabet=["a", "b"],
...     input_shapes=[{"features": [torch.Size([1, 8, 1, 4])]}],
... )
>>> predictions = torch.zeros(1, 1, 3)
>>> target = torch.tensor([[ord("a")]])
>>> loss = CTCLoss(use_focal_loss=False, node=head)
>>> round(loss(predictions, target).item(), 4)
1.0986
>>> focal_loss = CTCLoss(node=head)
>>> round(focal_loss(predictions, target).item(), 4)
0.4883
```

Parameters

 * `predictions` (`Tensor`): Logits of shape `[B, T, C]`, the main output of the node. `T` is the sequence length, and `C` is the
   number of classes, with class `0` as blank.
 * `target` (`Tensor`): The `metadata/text` label of shape `[B, L]`. Each value is the Unicode code point of one character, and
   `0` pads the shorter texts.

Returns

 * `Tensor`: The mean loss over the batch, as a scalar.

#### Attributes

##### loss_func

##### node
