Segformer¶
tiatoolbox.models.architecture.segformer.Segformer
- class Segformer(encoder_name='mit_b5', decoder_segmentation_channels=256, in_channels=3, classes=1, activation=None, upsampling=4)[source]¶
SegFormer semantic segmentation model (Hugging Face Transformers).
- Parameters:
encoder_name (str) – Mix Transformer backbone name:
mit_b0…mit_b5.decoder_segmentation_channels (int) – Channel width of the all-MLP decoder (HF
decoder_hidden_size).in_channels (int) – Number of input image channels (default RGB = 3).
classes (int) – Number of output segmentation classes (
num_labels).activation (nn.Module | None) – Optional activation applied after the classifier / upsample.
upsampling (int) – Spatial upsample factor applied to HF logits (HF outputs at 1/4 resolution; default
4restores input size).
Initializes the SegFormer model.
Methods
Run encoder-decoder and upsample logits to the input resolution.
Run inference on a batch of images.
Load HF or SMP SegFormer weights.
Postprocess model output to generate a class mask.
Preprocess input image for inference.
Attributes
training- forward(x, *args, **kwargs)[source]¶
Run encoder-decoder and upsample logits to the input resolution.
- static infer_batch(model, batch_data, *, device)[source]¶
Run inference on a batch of images.
Transfers the model and input batch to the specified device, performs forward pass, and returns softmax probabilities.
- Parameters:
model (Segformer) – Segformer model instance.
batch_data (torch.Tensor) – Batch of input images in NHWC format.
device (str) – Device for inference (e.g., “cpu” or “cuda”).
- Returns:
Inference results as a NumPy array of shape (N, H, W, C).
- Return type:
np.ndarray
Example
>>> batch = torch.randn(4, 448, 448, 3) >>> probs = Segformer.infer_batch( ... model, batch, device="cpu" ... ) >>> probs.shape (4, 448, 448, 1)
- load_state_dict(state_dict, strict=True, assign=False)[source]¶
Load HF or SMP SegFormer weights.
SMP checkpoints (keys like
encoder.patch_embed1) are remapped to Hugging Face names. Classifier tensors keep their shapes, so checkpoints that only differ by output class count load whenclassesmatches.- Parameters:
state_dict (Mapping[str, torch.Tensor])
strict (bool)
assign (bool)
- Return type:
_IncompatibleKeys
- postproc(image)[source]¶
Postprocess model output to generate a class mask.
Applies argmax and morphological operations to classify pixels.
- Parameters:
image (np.ndarray) – Input probability map as a NumPy array of shape (H, W, C).
self (Segformer)
- Returns:
Tissue mask
- Return type:
np.ndarray
Example
>>> model = Segformer(classes=2) >>> mask = model.postproc(probs) >>> mask.shape (448, 448)
- static preproc(image)[source]¶
Preprocess input image for inference.
Applies ImageNet normalization to the input image.
- Parameters:
image (np.ndarray) – Input image as a NumPy array of shape (H, W, C) in uint8 format.
- Returns:
Preprocessed image normalized to ImageNet statistics.
- Return type:
np.ndarray
Example
>>> img = np.random.randint(0, 255, (448, 448, 3), dtype=np.uint8) >>> processed = Segformer.preproc(img) >>> processed.shape (448, 448, 3)