# How to Fine-Tune SAM for Specific Tasks: A Complete Technical Guide

> Learn how to fine-tune the Segment Anything Model SAM for your specific tasks. This guide details efficient methods including freezing the image encoder or full model adaptation for optimal results.

- Repository: [Meta Research/segment-anything](https://github.com/facebookresearch/segment-anything)
- Tags: how-to-guide
- Published: 2026-03-07

---

**Yes, you can fine-tune the Segment Anything Model (SAM) for specific tasks by freezing the heavy image encoder and training the lightweight prompt encoder and mask decoder, or by full-model adaptation depending on your dataset size.**

The `facebookresearch/segment-anything` repository provides a modular transformer-based architecture that separates visual encoding from prompt encoding and mask decoding. This decoupling makes it straightforward to adapt SAM to downstream segmentation tasks without requiring massive computational resources.

## Understanding SAM's Modular Architecture

SAM is built from three distinct components that can be independently frozen or trained. Understanding these modules is essential for designing an efficient fine-tuning strategy.

### Image Encoder (Vision Transformer)

The **image encoder** is a Vision Transformer (ViT) backbone that processes padded `1024×1024` RGB images into dense embedding maps. Implemented in [`segment_anything/modeling/image_encoder.py`](https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/modeling/image_encoder.py) as `ImageEncoderViT`, this component contains the majority of SAM's parameters (approximately 632M for ViT-H). Freezing this module during fine-tuning preserves generic visual features learned from the SA-1B dataset while reducing memory requirements.

### Prompt Encoder

The **prompt encoder**, located in [`segment_anything/modeling/prompt_encoder.py`](https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/modeling/prompt_encoder.py), converts user inputs—points, bounding boxes, and mask prompts—into sparse token embeddings and dense feature maps. This lightweight module (only ~200K parameters) is ideal for adaptation when your task involves domain-specific prompting strategies, such as medical imaging points or custom bounding box formats.

### Mask Decoder

The **mask decoder** in [`segment_anything/modeling/mask_decoder.py`](https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/modeling/mask_decoder.py) uses a two-way transformer (depth=2) to attend to image embeddings and prompt tokens, producing mask logits and IoU quality scores. With approximately 4M parameters, this module can be quickly fine-tuned to learn task-specific mask boundaries while maintaining the pretrained attention mechanisms.

| Component | File Path | Parameters (ViT-H) | Primary Role |
|-----------|-----------|-------------------|--------------|
| Image Encoder | [`segment_anything/modeling/image_encoder.py`](https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/modeling/image_encoder.py) | ~632M | Dense visual features |
| Prompt Encoder | [`segment_anything/modeling/prompt_encoder.py`](https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/modeling/prompt_encoder.py) | ~200K | Sparse/dense prompt embeddings |
| Mask Decoder | [`segment_anything/modeling/mask_decoder.py`](https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/modeling/mask_decoder.py) | ~4M | Mask prediction & IoU scoring |
| SAM Wrapper | [`segment_anything/modeling/sam.py`](https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/modeling/sam.py) | - | Orchestration & I/O handling |

## Fine-Tuning Strategies for SAM

The optimal fine-tuning approach depends on your dataset size, computational budget, and how much your target domain diverges from natural images. SAM's modular design supports four primary strategies.

### Feature-Extractor Approach

Freeze the `image_encoder` and train only the `prompt_encoder` and `mask_decoder`. This is the most memory-efficient strategy, requiring only ~4.2M trainable parameters. Use this approach when your dataset is small (<10K images) or when your domain (e.g., medical, satellite) still shares low-level visual features with natural images.

### Full-Model Fine-Tuning

Unfreeze all modules including the ViT backbone. This requires significant GPU memory (16GB+ for ViT-H with gradient checkpointing) but yields the best performance when your target domain differs substantially from the SA-1B training distribution, such as synthetic aperture radar imagery or electron microscopy.

### Adapter Layers

Insert tiny bottleneck MLPs (adapters) after the image encoder blocks and train only these new parameters (~1-2M additional parameters). This preserves the original pretrained weights while learning domain shifts, making it ideal for multi-task scenarios where you need to switch between domains without catastrophic forgetting.

### Prompt-Only Adaptation

Freeze both encoders and the decoder, adding only a new classification head on top of the `iou_predictions` output. This strategy treats SAM as a fixed feature extractor for mask quality scoring, useful for tasks requiring binary mask acceptance/rejection based on learned quality metrics.

## Implementation: Loading and Preparing SAM for Training

Before implementing a training loop, you must load the pretrained weights and configure gradient computation for your chosen strategy. The `sam_model_registry` in [`segment_anything/build_sam.py`](https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/build_sam.py) provides convenient factory functions for all three model sizes (ViT-H, ViT-L, ViT-B).

### Loading and Freezing the Image Encoder

```python
from segment_anything import sam_model_registry
import torch

# Load pretrained SAM (ViT-Huge by default)

sam = sam_model_registry["vit_h"](checkpoint="sam_vit_h_4b8939.pth")
sam.train()  # Enable dropout and batch norm updates if present

# Freeze the image encoder to preserve generic visual features

for name, param in sam.named_parameters():
    if name.startswith("image_encoder"):
        param.requires_grad = False

# Verify trainable parameters

trainable_params = sum(p.numel() for p in sam.parameters() if p.requires_grad)
print(f"Trainable parameters: {trainable_params:,}")  # ~4.2M for prompt+decoder only

```

This configuration keeps the ~632M parameter ViT backbone frozen while allowing gradients to flow through the prompt encoder and mask decoder, reducing VRAM requirements by approximately 95% compared to full-model training.

## Writing the Training Loop

SAM expects input dictionaries containing preprocessed images and optional prompts. The forward pass returns mask logits and IoU predictions. When fine-tuning, you typically disable the multimask output to simplify loss computation, producing a single mask per prompt.

### Complete Training Loop with Dice Loss

```python
import torch
import torch.nn.functional as F
from segment_anything.utils.transforms import ResizeLongestSide

# Initialize optimizer for trainable parameters only

optimizer = torch.optim.Adam(
    filter(lambda p: p.requires_grad, sam.parameters()),
    lr=1e-4,
    weight_decay=1e-4
)

# Resize transform matches SAM's expected 1024x1024 input

resize_transform = ResizeLongestSide(sam.image_encoder.img_size)

def dice_loss(pred, target, smooth=1e-6):
    """Compute Dice loss for binary segmentation."""
    pred = torch.sigmoid(pred)
    intersection = (pred * target).sum(dim=(2, 3))
    union = pred.sum(dim=(2, 3)) + target.sum(dim=(2, 3))
    dice = (2. * intersection + smooth) / (union + smooth)
    return 1 - dice.mean()

# Training loop

num_epochs = 10
for epoch in range(num_epochs):
    for batch in train_loader:
        # batch contains: "image", "original_size", "point_coords", 

        # "point_labels", "ground_truth_mask"

        
        # Forward pass - single mask output for simplicity

        outputs = sam([batch], multimask_output=False)[0]
        
        pred_masks = outputs["masks"]  # [B, 1, H, W]

        iou_preds = outputs["iou_predictions"]  # [B, 1]

        
        # Compute loss against ground truth

        gt_masks = batch["ground_truth_mask"].unsqueeze(1).float()
        
        # Resize GT to match prediction size if necessary

        if gt_masks.shape[-2:] != pred_masks.shape[-2:]:
            gt_masks = F.interpolate(
                gt_masks, 
                size=pred_masks.shape[-2:], 
                mode="bilinear", 
                align_corners=False
            )
        
        loss = dice_loss(pred_masks, gt_masks)
        
        # Optional: Add IoU prediction loss (MSE against actual IoU)

        with torch.no_grad():
            pred_binary = (torch.sigmoid(pred_masks) > 0.5).float()
            intersection = (pred_binary * gt_masks).sum(dim=(2, 3))
            union = (pred_binary + gt_masks).clamp(0, 1).sum(dim=(2, 3))
            true_iou = (intersection / (union + 1e-6)).squeeze(1)
        
        iou_loss = F.mse_loss(iou_preds.squeeze(1), true_iou)
        total_loss = loss + 0.1 * iou_loss
        
        optimizer.zero_grad()
        total_loss.backward()
        optimizer.step()

```

The forward pass expects a list of dictionaries matching the format described in [`segment_anything/modeling/sam.py`](https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/modeling/sam.py) (lines 58-84). Each dictionary must contain the preprocessed image tensor and original size, with optional prompt coordinates and labels.

## Extending SAM with Custom Heads

For tasks requiring mask quality classification or domain-specific scoring, you can freeze SAM's core components and attach lightweight prediction heads to the existing IoU prediction features.

### Adding a Classification Head on IoU Predictions

```python
import torch.nn as nn

class SamWithClassifier(nn.Module):
    def __init__(self, sam_model, num_classes):
        super().__init__()
        self.sam = sam_model
        
        # Extract feature dimension from IoU head's final layer

        iou_features = sam_model.mask_decoder.iou_prediction_head.layers[-1].out_features
        
        # New classification head

        self.classifier = nn.Linear(iou_features, num_classes)
        
    def forward(self, batch, multimask_output=False):
        # Get SAM outputs

        outputs = self.sam([batch], multimask_output=multimask_output)[0]
        
        # Extract IoU features for classification

        iou_features = outputs["iou_predictions"]  # [B, num_masks]

        class_logits = self.classifier(iou_features)
        
        return class_logits, outputs["masks"]

```

This approach leverages the mask decoder's existing IoU prediction head (`iou_prediction_head` in [`segment_anything/modeling/mask_decoder.py`](https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/modeling/mask_decoder.py)) to generate features for downstream classification, keeping the segmentation backbone frozen.

## Exporting Fine-Tuned Models for Production

After fine-tuning, you can export SAM to ONNX format for deployment in production environments. The repository includes a dedicated export script that handles the mask decoder conversion.

### Exporting to ONNX

```bash
python scripts/export_onnx_model.py \
    --checkpoint my_finetuned_sam.pth \
    --model-type vit_h \
    --output sam_finetuned.onnx \
    --return-single-mask

```

The export script ([`scripts/export_onnx_model.py`](https://github.com/facebookresearch/segment-anything/blob/main/scripts/export_onnx_model.py)) converts the mask decoder to ONNX while preserving the image encoder's output format. Use the `--return-single-mask` flag when you've fine-tuned for single-mask prediction to optimize inference speed.

## Summary

- **SAM's modular architecture** separates the heavy ViT image encoder ([`segment_anything/modeling/image_encoder.py`](https://github.com/facebookresearch/segment-anything/blob/main/segment_anything/modeling/image_encoder.py)) from the lightweight prompt encoder and mask decoder, enabling targeted fine-tuning with minimal compute.
- **Freeze the image encoder** to reduce trainable parameters by ~95% while adapting the ~4M parameter mask decoder and prompt encoder to your domain.
- **Use standard PyTorch training loops** with the `sam_model_registry` loader and dictionaries containing `"image"`, `"original_size"`, and prompt coordinates as expected by `Sam.forward()`.
- **Extend functionality** by attaching custom heads to the IoU prediction features or exporting fine-tuned weights to ONNX using [`scripts/export_onnx_model.py`](https://github.com/facebookresearch/segment-anything/blob/main/scripts/export_onnx_model.py).

## Frequently Asked Questions

### Can I fine-tune SAM on a small dataset?

Yes. Freeze the image encoder (`sam.image_encoder`) and train only the prompt encoder and mask decoder. This reduces trainable parameters from ~636M to ~4.2M, preventing overfitting on datasets with fewer than 10,000 images while still adapting mask generation to your domain.

### Which SAM components should I freeze to save memory?

Freeze `sam.image_encoder` (the ViT backbone) to save approximately 95% of GPU memory during training. For extreme memory constraints, also freeze `sam.prompt_encoder` and train only the `sam.mask_decoder`, requiring less than 4GB VRAM for batch size 1 at 1024×1024 resolution.

### How do I export my fine-tuned SAM model to ONNX?

Use the provided export script: `python scripts/export_onnx_model.py --checkpoint my_finetuned_sam.pth --model-type vit_h --output model.onnx`. Ensure your checkpoint contains the full model state dict, and use the `--return-single-mask` flag if you fine-tuned for single-mask output to optimize the exported graph.

### What loss function should I use for fine-tuning SAM?

Use **Dice loss** for mask prediction to handle class imbalance in segmentation, combined with **MSE loss** for the IoU prediction head. Calculate the true IoU between predicted and ground-truth masks during training to supervise the `iou_predictions` output, helping SAM learn domain-specific mask quality estimation.