# classification_head

Python API: `luxonis_train.nodes.heads.classification_head`

A classification head with pooling, dropout, and a linear layer.

## Classes

### ClassificationHead

Classification head with one linear layer.

 * `Inputs:`: * `inputs` (`Tensor`): [B, C, h, w]
 * `Outputs:`: * `classification` (`Tensor`): [B, nclasses] logits

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

> **Notes**
> The head applies global average pooling, dropout, and one linear layer. The feature map can have any height h and width w. `forward` does not check the mode, so export mode also gives the logits. The dropout acts only in training mode.

 * `Variants:`: None. Configure the node through `params`.

> **Example**
> A node entry in the `model.nodes` section of a config:

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

 * `Compatible with:`: * Attach index: `-1`, the last output of the input node
    * Required labels: `classification`
    * 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)
    * Losses: *
      [BCEWithLogitsLoss](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/bce_with_logits.md)
       * [CrossEntropyLoss](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/cross_entropy.md)
       * [OHEMLoss](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/ohem_loss.md)
       * [SigmoidFocalLoss](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/sigmoid_focal_loss.md)
       * [SmoothBCEWithLogitsLoss](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/losses/smooth_bce_with_logits.md)
       * [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)
    * Metrics: *
      [Accuracy](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md)
       * [ConfusionMatrix](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/confusion_matrix/confusion_matrix.md)
       * [F1Score](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md)
       * [JaccardIndex](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md)
       * [Precision](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md)
       * [Recall](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/torchmetrics.md)
    * Visualizers:
      [ClassificationVisualizer](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/visualizers/classification_visualizer.md)
    * Export parser: `ClassificationParser`

#### Methods

##### init

```python
def __init__(dropout_rate: float = 0.2, **kwargs):
```

Build the pooling, the dropout, and the linear layer.

Parameters

 * `dropout_rate` (`float`): The probability that the dropout layer sets a pooled feature to zero in training mode, in `[0, 1]`.
 * `**kwargs`: Keyword arguments for
   [BaseNode](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md).
   They must hold `input_shapes` or `in_sizes`, and the class count through `n_classes` or `dataset_metadata`.

##### forward

```python
def forward(inputs: Tensor) -> Tensor:
```

Compute the class logits.

> **Example**
> ```pycon
>>> import torch
>>> from torch import Size
>>> from luxonis_train.nodes import ClassificationHead
>>> head = ClassificationHead(
...     n_classes=4,
...     input_shapes=[{"features": [Size([2, 16, 7, 7])]}],
... )
>>> head(torch.zeros(2, 16, 7, 7)).shape
torch.Size([2, 4])
```

Parameters

 * `inputs` (`Tensor`): The feature map of shape `[B, C, h, w]`.

Returns

 * `Tensor`: The logits of shape `[B, n_classes]`. [BaseNode.run](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md) puts them under the `"classification"` key.

##### get_custom_head_config

```python
def get_custom_head_config(self) -> Params:
```

Return the head-specific metadata for the NN Archive.

A subclass overrides the method to give its parser more values. [get_head_config](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/base_head.md) merges the result into the `"metadata"` dictionary. The base implementation returns an empty dictionary.

Returns

 * `Params`: The additional metadata keys and their values.

#### Attributes

##### head

##### in_channels

The number of channels of the attached inputs.

It is the third dimension from the end of [in_sizes](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md), so a shape with or without the batch dimension gives the same value. A list of sizes gives a list of channel counts.

Raises

 * `RuntimeError`: When [in_sizes](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md) cannot find the input sizes.
 * `ValueError`: When [attach_index](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md) does not fit the sizes.

##### parser
