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
| Requirement | Value |
|---|---|
| Python | 3.11 or 3.12 |
| Platform | Linux, macOS, or Windows |
| Dataset | COCO annotations with an instance mask for every object |
| Base weights | Hugging Face access or a local SAM3 checkpoint |
| GPU | CUDA GPU strongly recommended |
Quick Installation
Install PyTorch for your CUDA version, then install Iris with the SAM3-LoRA dependencies:
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
"""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:
python train.py \
--dataset-dir path/to/dataset \
--output-dir results/sam3-loraResume a run from a saved checkpoint:
python train.py \
--dataset-dir path/to/dataset \
--output-dir results/sam3-lora \
--resume-from results/sam3-lora/latest.ptTraining 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:
cd telekinesis-examples
python examples/iris/train_custom_sam3_lora_model.py \
--dataset-dir path/to/dataset \
--output-dir results/sam3-loraAdd --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:
| Output | Description |
|---|---|
latest.pt | Updated after every completed epoch |
best.pt | Updated when the monitored validation or training loss improves |
epoch_NNNN.pt | Written according to checkpoint_interval |
| TensorBoard logs | Training 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
| Option | Default | Description |
|---|---|---|
class_names | None | Category-ID-to-prompt mapping or ordered prompt sequence; required by SAM3LoRA |
sam3_checkpoint_path | None | Optional local SAM3 base checkpoint |
lora_weights_path | None | Optional LoRA weights loaded after model construction |
load_from_hf | True | Loads the base SAM3 checkpoint from Hugging Face |
rank | 8 | Positive LoRA decomposition rank |
alpha | 16.0 | Positive LoRA scaling factor |
dropout | 0.1 | LoRA dropout probability in [0, 1) |
target_modules | standard projection names | Module-name suffixes eligible for LoRA injection |
apply_to_vision_encoder | True | Adapts the vision encoder |
apply_to_text_encoder | True | Adapts the text encoder |
apply_to_geometry_encoder | False | Adapts the geometry encoder |
apply_to_detr_encoder | True | Adapts the DETR encoder |
apply_to_detr_decoder | True | Adapts the DETR decoder |
apply_to_mask_decoder | False | Adapts the mask decoder |
resolution | 1008 | Positive square training resolution |
num_negative_prompts | 2 | Negative category prompts sampled per image |
SAM3LoRA Properties
| Property | Type | Description |
|---|---|---|
config | SAM3LoRAConfig | Public unified configuration |
model_config | SAM3LoRAConfig | Model-construction configuration |
adapter_config | SAM3LoRAConfig | Configuration used by the training adapter |
model | nn.Module | Underlying SAM3 LoRA PyTorch model |
adapter | SAM3LoRATrainingAdapter | Cached prompted-training adapter |
name | str | Always "sam3-lora" |
train_transforms | None | Preprocessing is handled by the adapter |
val_transforms | None | Preprocessing is handled by the adapter |
requires_masks | bool | Always 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:
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
rankwhen the default adapters do not have enough capacity. Higher ranks add trainable parameters and consume more memory. - Tune
alphawithrank; a common starting relationship isalphaequal to one or two times the rank. - Increase
dropoutwhen 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.
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.
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=1008initially. Higher resolutions increase memory and compute, while lower resolutions may lose small-object detail. - Reduce the Trainer
batch_sizebefore reducing resolution when GPU memory is limited. - Increase
num_negative_promptswhen 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_pathto 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.ptfor deployment selection andlatest.ptfor continuing interrupted training.