# torchmetrics

Python API: `luxonis_train.attached_modules.metrics.torchmetrics`

Metrics that wrap the classification metrics of `torchmetrics`.

[TorchMetricWrapper](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md)
infers `task` and the number of classes from the node that the metric attaches to, when the config does not set them. The other
arguments go to the `torchmetrics` metric, such as `average`.

## Classes

### Accuracy

Accuracy metric that wraps `torchmetrics.Accuracy`.

 * `Inputs:`: * `predictions` (`Tensor`): `[B, n_classes, ...]` logits or probabilities
    * `target` (`Tensor`): `[B, n_classes, ...]` one-hot or multi-hot labels
 * `Outputs:`: * `Tensor`: scalar
    * `tuple[Tensor, dict[str, Tensor]]`: the mean and the value of each class, when the built metric returns one value for each
      class, for example with `average: "none"`
 * `Formula:`: For `"binary"`, with the counts of true and false positives and negatives: Accuracy = (TP + TN)/(TP + TN + FP + FN)
   For `"multiclass"` and `"multilabel"`, `torchmetrics.Accuracy` counts the statistics of each class and combines the classes as
   its `average` argument selects.

> **References**
> * Source: Wraps [torchmetrics](https://github.com/Lightning-AI/torchmetrics) (Apache-2.0).
 * License: Apache-2.0 (this project)

> **Notes**
> [TorchMetricWrapper](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md) resolves `task` and the number of classes. The other `params` go to `torchmetrics.Accuracy`, for example `average`, `threshold`, or `top_k`.

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

```yaml
- name: ClassificationHead
  inputs: [ResNet]
  metrics:
    - name: Accuracy
```

 * `Compatible with:`: * Used by:
   [ClassificationModel](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/config/predefined_models/classification/v1/model.md)
    * Nodes: *
      [BiSeNetHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/bisenet_head.md)
       * [ClassificationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/classification_head.md)
       * [DDRNetSegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/ddrnet_segmentation_head.md)
       * [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)
       * [SegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/segmentation_head.md)
       * [TransformerClassificationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/transformer_classification_head.md)
       * [TransformerSegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/transformer_segmentation_head.md)

#### Attributes

##### supported_tasks

### F1Score

F1 score metric that wraps `torchmetrics.F1Score`.

 * `Inputs:`: * `predictions` (`Tensor`): `[B, n_classes, ...]` logits or probabilities
    * `target` (`Tensor`): `[B, n_classes, ...]` one-hot or multi-hot labels
 * `Outputs:`: * `Tensor`: scalar
    * `tuple[Tensor, dict[str, Tensor]]`: the mean and the value of each class, when the built metric returns one value for each
      class, for example with `average: "none"`
 * `Formula:`: With the counts of true positives, false positives, and false negatives: F1 = (2 TP)/(2 TP + FP + FN) For
   `"multiclass"` and `"multilabel"`, `torchmetrics.F1Score` counts the statistics of each class and combines the classes as its
   `average` argument selects.

> **References**
> * Source: Wraps [torchmetrics](https://github.com/Lightning-AI/torchmetrics) (Apache-2.0).
 * License: Apache-2.0 (this project)

> **Notes**
> [TorchMetricWrapper](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md) resolves `task` and the number of classes. The other `params` go to `torchmetrics.F1Score`, for example `average`, `threshold`, or `top_k`.

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

```yaml
- name: ClassificationHead
  inputs: [ResNet]
  metrics:
    - name: F1Score
```

 * `Compatible with:`: * Used by: *
   [ClassificationModel](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/config/predefined_models/classification/v1/model.md)
       * [SegmentationModel](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/config/predefined_models/segmentation/v1/model.md)
    * Nodes: *
      [BiSeNetHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/bisenet_head.md)
       * [ClassificationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/classification_head.md)
       * [DDRNetSegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/ddrnet_segmentation_head.md)
       * [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)
       * [SegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/segmentation_head.md)
       * [TransformerClassificationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/transformer_classification_head.md)
       * [TransformerSegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/transformer_segmentation_head.md)

#### Attributes

##### supported_tasks

### JaccardIndex

Jaccard index metric that wraps `torchmetrics.JaccardIndex`.

 * `Inputs:`: * `predictions` (`Tensor`): `[B, n_classes, ...]` logits or probabilities
    * `target` (`Tensor`): `[B, n_classes, ...]` one-hot or multi-hot labels
 * `Outputs:`: * `Tensor`: scalar
    * `tuple[Tensor, dict[str, Tensor]]`: the mean and the value of each class, when the built metric returns one value for each
      class, for example with `average: "none"`
 * `Formula:`: With the counts of true positives, false positives, and false negatives: J = (TP)/(TP + FP + FN) For `"multiclass"`
   and `"multilabel"`, `torchmetrics.JaccardIndex` counts the statistics of each class and combines the classes as its `average`
   argument selects.

> **References**
> * Source: Wraps [torchmetrics](https://github.com/Lightning-AI/torchmetrics) (Apache-2.0).
 * License: Apache-2.0 (this project)

> **Notes**
> [TorchMetricWrapper](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md) resolves `task` and the number of classes. The other `params` go to `torchmetrics.JaccardIndex`, for example `average`, `threshold`, or `ignore_index`.

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

```yaml
- name: ClassificationHead
  inputs: [ResNet]
  metrics:
    - name: JaccardIndex
```

 * `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)
       * [SegmentationModel](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/config/predefined_models/segmentation/v1/model.md)
    * Nodes: *
      [BiSeNetHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/bisenet_head.md)
       * [ClassificationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/classification_head.md)
       * [DDRNetSegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/ddrnet_segmentation_head.md)
       * [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)
       * [SegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/segmentation_head.md)
       * [TransformerClassificationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/transformer_classification_head.md)
       * [TransformerSegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/transformer_segmentation_head.md)

#### Attributes

##### supported_tasks

### Precision

Precision metric that wraps `torchmetrics.Precision`.

 * `Inputs:`: * `predictions` (`Tensor`): `[B, n_classes, ...]` logits or probabilities
    * `target` (`Tensor`): `[B, n_classes, ...]` one-hot or multi-hot labels
 * `Outputs:`: * `Tensor`: scalar
    * `tuple[Tensor, dict[str, Tensor]]`: the mean and the value of each class, when the built metric returns one value for each
      class, for example with `average: "none"`
 * `Formula:`: With the counts of true positives and false positives: Precision = (TP)/(TP + FP) For `"multiclass"` and
   `"multilabel"`, `torchmetrics.Precision` counts the statistics of each class and combines the classes as its `average` argument
   selects.

> **References**
> * Source: Wraps [torchmetrics](https://github.com/Lightning-AI/torchmetrics) (Apache-2.0).
 * License: Apache-2.0 (this project)

> **Notes**
> [TorchMetricWrapper](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md) resolves `task` and the number of classes. The other `params` go to `torchmetrics.Precision`, for example `average`, `threshold`, or `top_k`.

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

```yaml
- name: ClassificationHead
  inputs: [ResNet]
  metrics:
    - name: Precision
```

 * `Compatible with:`: * Nodes: *
   [BiSeNetHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/bisenet_head.md)
       * [ClassificationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/classification_head.md)
       * [DDRNetSegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/ddrnet_segmentation_head.md)
       * [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)
       * [SegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/segmentation_head.md)
       * [TransformerClassificationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/transformer_classification_head.md)
       * [TransformerSegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/transformer_segmentation_head.md)

#### Attributes

##### supported_tasks

### Recall

Recall metric that wraps `torchmetrics.Recall`.

 * `Inputs:`: * `predictions` (`Tensor`): `[B, n_classes, ...]` logits or probabilities
    * `target` (`Tensor`): `[B, n_classes, ...]` one-hot or multi-hot labels
 * `Outputs:`: * `Tensor`: scalar
    * `tuple[Tensor, dict[str, Tensor]]`: the mean and the value of each class, when the built metric returns one value for each
      class, for example with `average: "none"`
 * `Formula:`: With the counts of true positives and false negatives: Recall = (TP)/(TP + FN) For `"multiclass"` and
   `"multilabel"`, `torchmetrics.Recall` counts the statistics of each class and combines the classes as its `average` argument
   selects.

> **References**
> * Source: Wraps [torchmetrics](https://github.com/Lightning-AI/torchmetrics) (Apache-2.0).
 * License: Apache-2.0 (this project)

> **Notes**
> [TorchMetricWrapper](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md) resolves `task` and the number of classes. The other `params` go to `torchmetrics.Recall`, for example `average`, `threshold`, or `top_k`.

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

```yaml
- name: ClassificationHead
  inputs: [ResNet]
  metrics:
    - name: Recall
```

 * `Compatible with:`: * Used by:
   [ClassificationModel](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/config/predefined_models/classification/v1/model.md)
    * Nodes: *
      [BiSeNetHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/bisenet_head.md)
       * [ClassificationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/classification_head.md)
       * [DDRNetSegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/ddrnet_segmentation_head.md)
       * [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)
       * [SegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/segmentation_head.md)
       * [TransformerClassificationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/transformer_classification_head.md)
       * [TransformerSegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/transformer_segmentation_head.md)

#### Attributes

##### supported_tasks

### TorchMetricWrapper

Base class for the metrics that wrap a `torchmetrics` metric.

A subclass sets the class attribute `Metric` to a task wrapper of `torchmetrics`, such as `torchmetrics.Accuracy`.
[init](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md)
resolves `task` and the number of classes, and stores the built metric in the `metric` attribute.
[update](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md),
[compute](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md),
and
[reset](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md)
call the built metric. This class does not set `Metric`, so a config that names it fails with `AttributeError`. Name a subclass
instead.

 * `Inputs:`: * `predictions` (`Tensor`): `[B, n_classes, ...]` logits or probabilities
    * `target` (`Tensor`): `[B, n_classes, ...]` one-hot or multi-hot labels
 * `Outputs:`: * `Tensor`: scalar
    * `tuple[Tensor, dict[str, Tensor]]`: the mean and the value of each class, when the built metric returns one value for each
      class, for example with `average: "none"`

> **References**
> * Source: Wraps [torchmetrics](https://github.com/Lightning-AI/torchmetrics) (Apache-2.0).
 * License: Apache-2.0 (this project)

> **Notes**
> `task` is `"binary"`, `"multiclass"`, or `"multilabel"`. When the `params` do not set it, the metric infers it and logs a warning. For `"multiclass"`, [update](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md) converts the one-hot target to class indices. On an `anomaly_detection` node, the metric reads the `segmentation` label.

> **Example**
> A subclass, here [Accuracy](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md), attached to a `DiscSubNetHead` in `model.nodes`:

```yaml
- name: DiscSubNetHead
  inputs: [RecSubNet]
  metrics:
    - name: Accuracy
```

 * `Compatible with:`: * Nodes: *
   [BiSeNetHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/bisenet_head.md)
       * [DDRNetSegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/ddrnet_segmentation_head.md)
       * [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)
       * [SegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/segmentation_head.md)
       * [TransformerSegmentationHead](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/transformer_segmentation_head.md)

#### Methods

##### init

```python
def __init__(**kwargs):
```

Resolve the task and build the wrapped `torchmetrics` metric.

`task` comes from `kwargs`. Without it, the method infers `"binary"` from `num_classes` of `1`, `"multiclass"` from another
`num_classes`, and `"multilabel"` from `num_labels`. Without these arguments, it infers `"binary"` for a node with one class and
`"multiclass"` for other nodes. It logs a warning when it infers the task.

The number of classes comes from `num_classes`, then from `num_labels`, then from the node. The method passes it to `Metric` as
`num_classes` for `"multiclass"`, and as `num_labels` for `"multilabel"`.

> **Example**
> ```pycon
>>> metric = Accuracy(task="multilabel", num_labels=3)
>>> type(metric.metric).__name__
'MultilabelAccuracy'
```

Parameters

 * `**kwargs`: `node` goes to [BaseMetric](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/base_metric.md). The other arguments and the resolved `task` go to the constructor of `Metric`, for example `average` or `threshold`.

Raises

 * `ValueError`: When the method cannot infer `task`, or when `task` is not `"binary"`, `"multiclass"`, or `"multilabel"`. Also when the number of classes does not fit the task: unknown or `1` for `"multiclass"` and `"multilabel"`, or more than `1` for `"binary"`.

##### compute

```python
def compute(self) -> Tensor | tuple[Tensor, dict[str, Tensor]]:
```

Compute the value of the wrapped metric.

The method returns a result with one element unchanged. It treats a result with more than one element as one value for each class, as with `average: "none"`. It then returns the mean of the values and a dictionary of the class values. The keys are `"<Metric>_<class name>"`, where `<Metric>` is the class name of the built metric, such as `MulticlassAccuracy`. The class names come from the node. Without a node, the `classes` property raises `RuntimeError`.

> **Example**
> One correct prediction out of two:

```pycon
>>> import torch
>>> metric = Accuracy(task="multiclass", num_classes=3)
>>> predictions = torch.tensor([[2.0, 0.0, 0.0], [0.0, 2.0, 0.0]])
>>> target = torch.tensor([[1.0, 0.0, 0.0], [0.0, 0.0, 1.0]])
>>> metric.update(predictions, target)
>>> metric.compute().item()
0.5
```

Returns

 * `Tensor | tuple[Tensor, dict[str, Tensor]]`: The scalar value, or the mean of the class values and the dictionary of the class values.

Raises

 * `ValueError`: When the result has more than one element and the node has no class names.

##### reset

```python
def reset(self):
```

Reset the states of the wrapped metric.

The method calls `reset` only on the `metric` attribute. It does not call the `reset` of the `torchmetrics` base class on the wrapper. Thus the wrapper keeps the cached result of its last [compute](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md), and [compute](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md) returns that result until the next [update](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md).

##### update

```python
def update(predictions: Tensor, target: Tensor):
```

Add one batch to the wrapped metric.

For `"multiclass"`, the method converts `target` to class indices with `argmax` over dimension `1`. For the other tasks, it passes `target` unchanged.

Parameters

 * `predictions` (`Tensor`): The main output of the node, of shape `[B, n_classes, ...]`. Logits or probabilities.
 * `target` (`Tensor`): The label of the task, of shape `[B, n_classes, ...]`. It is a binary mask for `"binary"`, one-hot for `"multiclass"`, and multi-hot for `"multilabel"`.

#### Attributes

##### Metric

##### metric

##### required_labels

The labels for the `target` parameter of [update](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md).

[BaseAttachedModule.get_parameters](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/base_attached_module.md) reads this set for the `target` parameter of [update](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md), which has no label suffix. On an `anomaly_detection` node, the set holds only `"segmentation"`, so `target` receives the anomaly mask. On other nodes, it holds the labels of the task. The property raises `RuntimeError` when the metric has no task.
