# transformer_classification_head

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

A classification head over the CLS token of a transformer backbone.

## Classes

### TransformerClassificationHead

Classification head for the CLS token of a transformer backbone.

[DinoV3](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/backbones/dinov3/dinov3.md)
gives the CLS token when its `return_sequence` param is `True`.

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

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

> **Notes**
> The head applies dropout and one linear layer to the CLS token. It has no pooling, because the CLS token has no spatial dimensions. The dropout acts only in training mode. The mode does not change the output key or shape. This includes export mode.

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

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

```yaml
- name: TransformerClassificationHead
  inputs: [DinoV3]
```

 * `Compatible with:`: * Attach index: `-1`, the last output of the input node
    * Required labels: `classification`
    * 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 dropout and the linear layer.

The linear layer maps
[in_channels](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/transformer_classification_head.md)
values to `n_classes` logits.

Parameters

 * `dropout_rate` (`float`): The probability that the dropout layer sets a value of the CLS token to zero in training mode, in
   `[0, 1]`. The layer scales the other values by 1 ⁄ (1 − p), where p is `dropout_rate`. Defaults to `0.2`.
 * `**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 the input sizes through `input_shapes` or `in_sizes`, and the class count through `n_classes` or
   `dataset_metadata`.

##### forward

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

Compute the class logits from the CLS token.

The method applies the dropout and then the linear layer. The dropout acts only in training mode.

> **Example**
> ```pycon
>>> import torch
>>> from torch import Size
>>> from luxonis_train.nodes import TransformerClassificationHead
>>> head = TransformerClassificationHead(
...     n_classes=5,
...     input_shapes=[{"features": [Size([2, 384])]}],
... )
>>> head.in_channels
384
>>> head(torch.zeros(2, 384)).shape
torch.Size([2, 5])
```

Parameters

 * `x` (`Tensor`): The CLS token of shape `[B, C]`, where `C` is the embedding size.

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

##### attach_index

##### dropout

##### fc

##### in_channels

The embedding size of the CLS token.

It is the last dimension of [BaseNode.in_sizes](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md), the `C` of the input shape `[B, C]`. It replaces [BaseNode.in_channels](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md), which reads the third dimension from the end.

Raises

 * `TypeError`: When [BaseNode.in_sizes](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md) is a list of sizes. This occurs for an `attach_index` of `"all"` or a range. The constructor reads the property, so it raises the error too.

##### parser
