# transformer_segmentation_head

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

A segmentation head over the patch tokens of a transformer backbone.

## Classes

### TransformerSegmentationHead

Semantic segmentation head for the feature maps of a transformer.

Section 6.3.2 of the [DINOv3 paper](https://arxiv.org/abs/2508.10104) puts a ViT-Adapter without the injection and a Mask2Former
decoder on the backbone. This head replaces Mask2Former with a small convolutional decoder.
[DinoV3](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/backbones/dinov3/dinov3.md)
gives `depth` feature maps of one resolution when its `return_sequence` param is `False`.

 * `Inputs:`: * `inputs` (`list[Tensor]`): [B, Ci, hi, wi] per map, any resolution
 * `Outputs:`: * `segmentation` (`Tensor`): [B, nclasses, H, W] logits

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

> **Notes**
> A `1x1` convolution, a batch norm, and a ReLU project each feature map to `256` channels. The head resizes each projected map to a quarter of the image height and width, and averages the maps. A `3x3` convolution with ReLU and a `1x1` convolution then give one logit map for each class. A last resize brings the logits to the image size. All resizes are bilinear. 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: TransformerSegmentationHead
  inputs: [DinoV3]
```

 * `Compatible with:`: * Attach index: `"all"`, every output of the input node
    * Required labels: `segmentation`
    * 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)
       * [DiceCoefficient](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/dice_coefficient.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)
       * [MIoU](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/metrics/mean_iou.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:
      [SegmentationVisualizer](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/visualizers/segmentation_visualizer.md)
    * Export parser: `SegmentationParser`

#### Methods

##### init

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

Build the decoder and one projection for each feature map.

The constructor reads the channel count of each map from index `1` of its size. Each size must thus have the form `[B, C_i, h_i,
w_i]`.

The class annotates
[BaseNode.in_sizes](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md)
as a list. An integer `attach_index` selects one output, so
[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 single size. Such an `attach_index` thus makes the constructor raise
[IncompatibleError](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/utils/exceptions.md).

Parameters

 * `**kwargs` (`Any`): 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](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/heads/transformer_segmentation_head.md)
   also needs `original_in_shape`.

##### forward

```python
def forward(x: list[Tensor]) -> Tensor:
```

Decode the feature maps into segmentation logits.

`H` and `W` are the image size from `original_in_shape`. The method does these steps:

 * It projects each feature map to `256` channels.
 * It resizes each projected map to `[H // 4, W // 4]`.
 * It averages the resized maps.
 * It applies the decoder, which gives `n_classes` channels.
 * It resizes the logits to `[H, W]`.

Both resizes use bilinear interpolation.

> **Example**
> ```pycon
>>> import torch
>>> from torch import Size
>>> from luxonis_train.nodes import TransformerSegmentationHead
>>> sizes = [Size([1, 24, 4, 4]), Size([1, 12, 8, 8])]
>>> head = TransformerSegmentationHead(
...     n_classes=3,
...     input_shapes=[{"features": sizes}],
...     original_in_shape=Size([3, 64, 64]),
... )
>>> head([torch.zeros(size) for size in sizes]).shape
torch.Size([1, 3, 64, 64])
```

Parameters

 * `x` (`list[Tensor]`): The feature maps, each of shape `[B, C_i, h_i, w_i]`, in the order of the input sizes of the constructor. The maps can have different sizes.

Returns

 * `Tensor`: The logits of shape `[B, n_classes, H, W]`. [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 `"segmentation"` key.

#### Attributes

##### head

##### in_sizes

The sizes of the attached inputs.

The property uses the first rule that applies:

 1. The `in_sizes` constructor argument, when it is set.
 2. The `"features"` entry of the only input packet.
 3. The only entry of that packet.
 4. The entries of that packet whose keys match the names of the `forward` parameters. Their sizes must be equal, and the property uses the first one.

Rules 2 to 4 pass the entry through [get_attached](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md). The result is a single size for an integer [attach_index](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md), and a list of sizes for `"all"` or a range. A node with more than one input, or with shapes that the rules do not fit, must read [input_shapes](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md) instead.

> **Example**
> ```pycon
>>> from torch import Size, Tensor
>>> from luxonis_train.nodes import BaseNode
>>> class Node(BaseNode, register=False):
...     def forward(self, x: list[Tensor]) -> list[Tensor]:
...         return x
>>> shapes = [
...     {"features": [Size([2, 8, 64, 64]), Size([2, 16, 32, 32])]}
... ]
>>> node = Node(input_shapes=shapes)
>>> node.attach_index
'all'
>>> node.in_sizes
[torch.Size([2, 8, 64, 64]), torch.Size([2, 16, 32, 32])]
>>> node.in_channels, node.in_height, node.in_width
([8, 16], [64, 32], [64, 32])
```

Raises

 * `RuntimeError`: When `input_shapes` is missing or does not hold exactly one packet. Also when no key matches a `forward`
   parameter, or when the matching sizes differ. Also when
   [attach_index](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md)
   is `None` and the entry is a list.
 * `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.

##### n_classes

The number of classes of the node task.

The `n_classes` constructor argument comes first. Without it, the value comes from
[dataset_metadata](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md)
for
[task_name](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md).

Raises

 * `RuntimeError`: When the constructor got neither `n_classes` nor `dataset_metadata`.
 * `ValueError`: When the dataset has no task named
   [task_name](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md).

##### parser

##### projections
