core/nn/vision/sam_mask_decoder library

SAM mask decoder — turns (image embedding + prompt embeddings) into segmentation masks + IoU predictions.

Ports segment_anything.modeling.mask_decoder.MaskDecoder and its TwoWayTransformer (two-way cross-attention between prompt tokens and image features). Given the outputs of SamImageEncoder and SamPromptEncoder, produces:

  • masks — [num_masks, H_out, W_out] where H_out = W_out = imageEmbedH · 4 (256×256 for SAM ViT-B). These are raw mask logits; downstream code thresholds them at 0.
  • iou_predictions — [num_masks] scalar quality estimates for each mask.

SAM produces 4 mask tokens per forward: index 0 is the "no-mask" / single-output token, indices 1..3 are the multi-mask candidates. Call MaskDecoderOutput.select on the result to grab either the single mask (multimaskOutput: false) or the three candidates (multimaskOutput: true).

Modules built here (all specific to SAM):

  • SamAttention — MHA with an optional downsample_rate that reduces the internal Q/K/V dim before projecting back to embedDim. SAM's cross-attention layers use downsample_rate = 2 (embedDim 256 → internal 128) to save compute.
  • SamMlpBlock — two-layer MLP with GELU between, matching SAM's MLPBlock.
  • SamMlp — thin MLP wrapper: Linear → ReLU → … → Linear with configurable depth (used for the mask-hypernetwork and IoU heads).
  • SamTwoWayAttentionBlock — one block of the two-way transformer: self-attn on queries + q→k cross-attn + MLP + k→q cross-attn.
  • SamTwoWayTransformer — 2 × SamTwoWayAttentionBlock + a final token-to-image attention + LayerNorm.
  • SamMaskDecoder — the top-level module: prepends the IoU + mask tokens to the sparse prompt embeddings, runs the two-way transformer, upsamples the image features with two ConvTranspose2ds, applies per-mask hypernet MLPs, and dot- products with the upsampled features to produce mask logits.

Classes

MaskDecoderOutput
Output of SamMaskDecoder.call.
SamAttention
SamMaskDecoder
SamMaskDecoderConfig
Config for SamMaskDecoder. Defaults match facebook/sam-vit-*.
SamMlp
SAM's MLP helper: numLayers stacked Linear layers with ReLU between (final layer has no activation). Optional sigmoidOutput.
SamMlpBlock
SamTwoWayAttentionBlock
The two-way attention block from SAM's mask decoder:
SamTwoWayTransformer