core/nn/vision/sam_image_encoder library

SAM ViT-B image encoder — windowed attention + relative positional bias + neck.

Ports the image encoder from segment-anything (Kirillov et al. 2023, facebookresearch/segment-anything). Takes a preprocessed 3×1024×1024 image (channels-last-per-pixel patchified layout, matching the rest of our ViT stack) and produces the [N, 256, 64, 64] image embedding that SAM's prompt encoder + mask decoder consume.

What's here (this session):

  • LayerNorm2d — per-pixel LayerNorm over the channel axis of an NCHW tensor, matching SAM's LayerNorm2d in the neck.
  • SamWindowedAttention — MHA with two special SAM knobs:
    • Windowed partitioning: with windowSize < inputSize, the [H, W, D] feature map is padded to a multiple of windowSize, partitioned into non-overlapping windows, and MHA runs per window. This cuts attention cost from O((H·W)²) to O(W² · numWindows) (SAM ViT-B: 4× speedup at 64² tokens).
    • Decomposed relative-positional bias: two learned tables relPosH[2·H − 1, headDim] and relPosW[2·W − 1, headDim] where for each query row i and key row j the bias Q[i] · relPosH[i-j+W-1] + Q[i] · relPosW[.] is added to the attention logits before softmax. Matches SAM's add_decomposed_rel_pos in segment_anything/modeling/image_encoder.py bit-for-bit.
  • SamViTBlock — pre-LN + SamWindowedAttention + residual, pre-LN + MLP (fc1 → GELU → fc2) + residual.
  • SamImageEncoder — patch embed (Conv2d(3, embedDim, k=patchSize, s=patchSize, bias)), learned 2-D positional embedding, N × SamViTBlock with SAM's global-attention layer schedule (indices [2, 5, 8, 11] for ViT-B), and a final neck (Conv2d(embed, 256, k=1) → LayerNorm2d → Conv2d(256, 256, k=3, p=1) → LayerNorm2d).

Not yet: HF safetensors loader, prompt encoder (Fourier positional encoding + point/box/mask embeddings), and mask decoder (two-way transformer + upsampling). All three land in a follow-up session; the primitives here (WindowedAttention, ConvTranspose2d, LayerNorm2d, positional embeddings) are the pieces they'll need.

Classes

LayerNorm2d
LayerNorm applied over the channel axis of an [N, C, H, W] tensor. Matches segment_anything.modeling.common.LayerNorm2d. Learnable per-channel gamma and beta.
SamImageEncoder
SamImageEncoderConfig
SamViTBlock
SamWindowedAttention
Attention block for SAM's image encoder. Accepts a 2-D flat [H·W, embedDim] token sequence (which represents an [H, W, D] feature map) and returns the same shape.