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
LayerNorm2din 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 ofwindowSize, partitioned into non-overlapping windows, and MHA runs per window. This cuts attention cost fromO((H·W)²)toO(W² · numWindows)(SAM ViT-B: 4× speedup at 64² tokens). - Decomposed relative-positional bias: two learned
tables
relPosH[2·H − 1, headDim]andrelPosW[2·W − 1, headDim]where for each query rowiand key rowjthe biasQ[i] · relPosH[i-j+W-1] + Q[i] · relPosW[.]is added to the attention logits before softmax. Matches SAM'sadd_decomposed_rel_posinsegment_anything/modeling/image_encoder.pybit-for-bit.
- Windowed partitioning: with
- 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. Matchessegment_anything.modeling.common.LayerNorm2d. Learnable per-channelgammaandbeta. - 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.