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]whereH_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_ratethat reduces the internal Q/K/V dim before projecting back toembedDim. SAM's cross-attention layers usedownsample_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 → … → Linearwith 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
MLPhelper:numLayersstacked Linear layers with ReLU between (final layer has no activation). OptionalsigmoidOutput. - SamMlpBlock
- SamTwoWayAttentionBlock
- The two-way attention block from SAM's mask decoder:
- SamTwoWayTransformer