Skip to content

Train a SAM3-LoRA Model ​

SUMMARY

Fine-tune SAM3 for text-prompted instance segmentation with compact LoRA adapters. Configure prompts and adapter targets with SAM3LoRAConfig, prepare a mask-annotated COCO dataset, and run or resume training with the shared Iris Trainer.

Python Version Support

Iris training supports Python 3.11 and 3.12 only. Create an environment with one of these versions before installing dependencies.

Requirements ​

RequirementValue
Python3.11 or 3.12
PlatformLinux, macOS, or Windows
DatasetCOCO annotations with an instance mask for every object
Base weightsHugging Face access or a local SAM3 checkpoint
GPUCUDA GPU strongly recommended

Quick Installation ​

Install PyTorch for your CUDA version, then install Iris with the SAM3-LoRA dependencies:

bash
pip install torch torchvision --index-url https://download.pytorch.org/whl/cu128
pip install "telekinesis-iris[sam3-lora]"

PyTorch Installation

For CPU, macOS, or another CUDA version, select the appropriate command in the PyTorch installation guide.

Training Workflow ​

The example below derives text prompts from COCO category names, configures LoRA targets, creates training and validation datasets, records TensorBoard metrics, and optionally resumes from a checkpoint.

Create train.py ​

python
"""Fine-tune the SAM3 image model with LoRA on a COCO dataset."""

import argparse
from pathlib import Path

import torch
from loguru import logger

from telekinesis.iris import dataset, logger as iris_logger, models, trainer


def train_custom_sam3_lora_model_example(args: argparse.Namespace) -> None:
    """Build SAM3 LoRA and train it with the shared Iris trainer."""
    seed = 42

    # ===================== Configure Model ================================
    dataset_dir = args.dataset_dir

    metadata_dir = (
        dataset_dir / "train"
        if (dataset_dir / "train").is_dir()
        else dataset_dir
    )
    categories = dataset.COCODataset(metadata_dir).categories

    model_config = models.SAM3LoRAConfig(
        class_names=categories,
        sam3_checkpoint_path=None,
        lora_weights_path=None,
        load_from_hf=True,
        rank=16,
        alpha=32.0,
        dropout=0.1,
        apply_to_vision_encoder=True,
        apply_to_text_encoder=False,
        apply_to_geometry_encoder=True,
        apply_to_detr_encoder=True,
        apply_to_detr_decoder=True,
        apply_to_mask_decoder=True,
        resolution=1008,
        num_negative_prompts=3,
    )
    model = models.SAM3LoRA(config=model_config)

    # ===================== Prepare Dataset ================================
    training_dataset, validation_dataset = dataset.prepare_coco_datasets(
        dataset_dir=dataset_dir,
        validation_split=0.3,
        seed=seed,
        train_transforms=model.train_transforms,
        val_transforms=model.val_transforms,
        include_masks=model.requires_masks,
    )

    # ===================== Train Model ====================================
    output_dir = args.output_dir
    device = "cuda" if torch.cuda.is_available() else "cpu"

    logger.info(
        f"Training SAM3 LoRA on {len(training_dataset)} images using {device}"
    )

    metric_logger = iris_logger.TensorBoardMetricLogger(
        output_dir / "tensorboard"
    )
    # Tune these values for your dataset size and available GPU memory.
    model_trainer = trainer.Trainer(
        model=model,
        training_dataset=training_dataset,
        validation_dataset=validation_dataset,
        output_dir=output_dir,
        epochs=100,
        batch_size=1,
        learning_rate=5e-5,
        weight_decay=0.01,
        device=device,
        num_workers=2,
        gradient_clip_norm=1.0,
        mixed_precision=True,
        evaluation_interval=1,
        checkpoint_interval=1,
        seed=seed,
        metric_logger=metric_logger,
    )

    try:
        if args.resume_from is None:
            model_trainer.train()
        else:
            model_trainer.resume(args.resume_from)
    except KeyboardInterrupt:
        logger.info("Training interrupted; existing checkpoints were preserved.")
    except Exception as error:
        logger.error(f"An error occurred during training: {error}")

    tensorboard_dir = (output_dir / "tensorboard").resolve()
    logger.info(f'TensorBoard: tensorboard --logdir "{tensorboard_dir}"')


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--dataset-dir",
        type=Path,
        required=True,
        help="Path to the COCO dataset directory.",
    )
    parser.add_argument(
        "--output-dir",
        type=Path,
        required=True,
        help="Directory in which to save training outputs.",
    )
    parser.add_argument(
        "--resume-from",
        type=Path,
        help="Optional checkpoint from which to resume training.",
    )
    train_custom_sam3_lora_model_example(parser.parse_args())

Launch Training ​

Start a new training run:

bash
python train.py \
  --dataset-dir path/to/dataset \
  --output-dir results/sam3-lora

Resume a run from a saved checkpoint:

bash
python train.py \
  --dataset-dir path/to/dataset \
  --output-dir results/sam3-lora \
  --resume-from results/sam3-lora/latest.pt

Training Guidance

See Best Practices: Trainer for guidance on tuning batch size, learning rate, epochs, gradient clipping, workers, evaluation, and checkpoints.

Segmentation Masks and Prompts

Every annotated object must contain a non-empty COCO polygon or RLE segmentation. Use stable, descriptive category names because Iris uses them as the training prompts and stores them in the deployment bundle.

Run the Example ​

A runnable version is available in the Telekinesis examples repository:

bash
cd telekinesis-examples
python examples/iris/train_custom_sam3_lora_model.py \
  --dataset-dir path/to/dataset \
  --output-dir results/sam3-lora

Add --resume-from results/sam3-lora/latest.pt to continue a previous run.

Checkpoints and Metrics ​

The training script saves the following outputs to the configured output_dir:

OutputDescription
latest.ptUpdated after every completed epoch
best.ptUpdated when the monitored validation or training loss improves
epoch_NNNN.ptWritten according to checkpoint_interval
TensorBoard logsTraining loss, learning rate, validation box/mask metrics, and previews

SAM3-LoRA checkpoints store the compact LoRA state plus the metadata required to reconstruct the model and its prompts.

Model Reference ​

SAM3LoRAConfig ​

OptionDefaultDescription
class_namesNoneCategory-ID-to-prompt mapping or ordered prompt sequence; required by SAM3LoRA
sam3_checkpoint_pathNoneOptional local SAM3 base checkpoint
lora_weights_pathNoneOptional LoRA weights loaded after model construction
load_from_hfTrueLoads the base SAM3 checkpoint from Hugging Face
rank8Positive LoRA decomposition rank
alpha16.0Positive LoRA scaling factor
dropout0.1LoRA dropout probability in [0, 1)
target_modulesstandard projection namesModule-name suffixes eligible for LoRA injection
apply_to_vision_encoderTrueAdapts the vision encoder
apply_to_text_encoderTrueAdapts the text encoder
apply_to_geometry_encoderFalseAdapts the geometry encoder
apply_to_detr_encoderTrueAdapts the DETR encoder
apply_to_detr_decoderTrueAdapts the DETR decoder
apply_to_mask_decoderFalseAdapts the mask decoder
resolution1008Positive square training resolution
num_negative_prompts2Negative category prompts sampled per image

SAM3LoRA Properties ​

PropertyTypeDescription
configSAM3LoRAConfigPublic unified configuration
model_configSAM3LoRAConfigModel-construction configuration
adapter_configSAM3LoRAConfigConfiguration used by the training adapter
modelnn.ModuleUnderlying SAM3 LoRA PyTorch model
adapterSAM3LoRATrainingAdapterCached prompted-training adapter
namestrAlways "sam3-lora"
train_transformsNonePreprocessing is handled by the adapter
val_transformsNonePreprocessing is handled by the adapter
requires_masksboolAlways True

Trainer Reference

See Configuration: Trainer Reference for all Trainer constructor options, defaults, and methods.

SAM3 Best Practices ​

Prompts and Categories ​

  • Use stable, descriptive COCO category names because they become the text prompts used during training.
  • Keep category IDs and names consistent between training and validation datasets.
  • Avoid renaming or reordering prompts when resuming a checkpoint.
  • Inspect prompt-specific validation results when classes are visually similar or semantically ambiguous.

Base Checkpoint Loading ​

With the default load_from_hf=True, Iris downloads the SAM3 base checkpoint from Hugging Face. Use a local checkpoint when training must run offline or when a specific base checkpoint must be reproduced:

python
config = SAM3LoRAConfig(
    class_names=categories,
    sam3_checkpoint_path="checkpoints/sam3.pt",
    load_from_hf=False,
)

Checkpoint Loading

sam3_checkpoint_path is required when load_from_hf=False. Use lora_weights_path only when initializing the model with existing LoRA parameters.

LoRA Capacity ​

rank, alpha, and dropout control LoRA capacity and regularization:

  • Increase rank when the default adapters do not have enough capacity. Higher ranks add trainable parameters and consume more memory.
  • Tune alpha with rank; a common starting relationship is alpha equal to one or two times the rank.
  • Increase dropout when the adapted model overfits. Reduce it when training is underfitting or unstable because too much signal is being dropped.
  • Change one setting at a time and compare validation mask metrics.
python
config = SAM3LoRAConfig(
    class_names=categories,
    rank=16,
    alpha=32.0,
    dropout=0.1,
)

LoRA Target Selection ​

Enable only the model components that need adaptation. Adapting more components increases capacity, memory use, and training time.

python
config = SAM3LoRAConfig(
    class_names=categories,
    rank=16,
    alpha=32.0,
    apply_to_vision_encoder=True,
    apply_to_text_encoder=False,
    apply_to_geometry_encoder=True,
    apply_to_detr_encoder=True,
    apply_to_detr_decoder=True,
    apply_to_mask_decoder=True,
)

Start with the defaults, then expand the targets only when validation results show that more adaptation is needed. At least one enabled component must contain a name matching target_modules.

The default target_modules are q_proj, k_proj, v_proj, out_proj, qkv, proj, fc1, fc2, c_fc, c_proj, linear1, and linear2. Override this tuple only when the target SAM3 implementation uses different module names or you intentionally want a narrower adapter set.

Resolution and Negative Prompts ​

  • Keep resolution=1008 initially. Higher resolutions increase memory and compute, while lower resolutions may lose small-object detail.
  • Reduce the Trainer batch_size before reducing resolution when GPU memory is limited.
  • Increase num_negative_prompts when the model confuses absent categories with foreground objects.
  • Too many negative prompts increase work per image and may dilute positive supervision; validate each change.

Existing LoRA Weights and Resume ​

  • Use lora_weights_path to initialize a new model from existing LoRA parameters.
  • Use Trainer.resume() when continuing an Iris training run because it also restores optimizer, scheduler, scaler, history, and epoch state.
  • Keep prompts, resolution, LoRA rank, enabled targets, and target module names compatible with the saved weights.
  • Preserve best.pt for deployment selection and latest.pt for continuing interrupted training.

Next Steps ​