Skip to content

Train and Resume a Model ​

SUMMARY

Train and fine-tune computer vision models using the shared Iris Trainer. Select a supported TrainableModel, prepare COCO datasets, configure training, and resume runs from saved checkpoints.

Training Workflow ​

Iris provides a common training workflow across supported model architectures. Each TrainableModel encapsulates its architecture-specific configuration and training behavior, while Trainer manages the training lifecycle.

The following steps demonstrate the workflow using RF-DETR, with configuration alternatives for SAM3-LoRA.

Load the Dataset ​

Start by loading the training dataset using COCODataset. Iris supports COCO-formatted annotations containing images, bounding boxes, category labels, and optional instance segmentation masks.

python
from telekinesis.iris.dataset import COCODataset

metadata = COCODataset("dataset/train")
category_ids = metadata.coco.getCatIds()
categories = {
    category["id"]: category["name"]
    for category in metadata.coco.loadCats(category_ids)
}

The metadata configures the selected model:

  • RF-DETR: Uses category IDs to determine the class count.
  • SAM3-LoRA: Uses category names as text prompts during training.

Dataset Preparation

See Prepare a Dataset for supported annotations and dataset structure.

Configure the Trainable Model ​

Choose a supported model architecture and instantiate its TrainableModel. The model provides the network, preprocessing, loss computation, evaluation behavior, and export metadata needed by Trainer.

RF-DETR ​

python
from telekinesis.iris.models import RFDETR

model = RFDETR(
    variant="seg-nano",
    num_classes=max(category_ids),
    pretrained=True,
)
TaskAvailable variants
Object detectionnano, small, medium, base, large
Instance segmentationseg-nano, seg-small, seg-medium, seg-large, seg-xlarge, seg-2xlarge

Use pretrained weights when fine-tuning on a custom dataset. Start with a smaller model when validating a new training pipeline.

SAM3-LoRA ​

python
from telekinesis.iris.models import SAM3LoRA

model = SAM3LoRA(
    class_names=categories,
    resolution=1008,
    num_negative_prompts=2,
)

SAM3-LoRA keeps the base model frozen and optimizes compact LoRA parameters using text prompts derived from COCO category names.

SAM3-LoRA Dependencies

Install the optional dependencies before using SAM3-LoRA:

bash
pip install "telekinesis-iris[sam3-lora]"

By default, Iris loads the base SAM3 weights from Hugging Face. You can specify a local checkpoint through SAM3LoRAConfig.

Model Configuration

See Configuration for architecture-specific settings, including model resolution and LoRA parameters.

Prepare Training and Validation Datasets ​

Initialize the datasets using the preprocessing and mask requirements exposed by the selected model wrapper. This structure works across supported model families; wrappers that handle preprocessing internally return None for their transforms.

python
train_dataset = COCODataset(
    "dataset/train",
    transforms=model.train_transforms,
    include_masks=model.requires_masks,
)
valid_dataset = COCODataset(
    "dataset/valid",
    transforms=model.val_transforms,
    include_masks=model.requires_masks,
)

Segmentation Datasets

When training an instance segmentation model, every object must include a valid segmentation polygon or RLE mask in its COCO annotations.

Instantiate the Trainer ​

Pass the configured model and datasets to the shared Iris Trainer. It manages training, validation, optimization, checkpointing, and metric reporting independently of the selected model family.

Adjust the parameters for the selected architecture, dataset size, and available GPU memory.

python
from telekinesis.iris.logger import TensorBoardMetricLogger
from telekinesis.iris.trainer import Trainer

trainer = Trainer(
    model=model,
    training_dataset=train_dataset,
    validation_dataset=valid_dataset,
    output_dir="results/training",
    epochs=20,
    batch_size=4,
    learning_rate=1e-4,
    metric_logger=TensorBoardMetricLogger(
        "results/training/tensorboard"
    ),
)

Only one trainer instance is required for a training run. See Configuration for all training parameters.

Train the Model ​

python
history = trainer.train()

The trainer executes the training and validation loop using the selected model's implementation.

Checkpoints ​

Iris saves the state required to continue a run, including the model, criterion, optimizer, scheduler, gradient scaler, metric history, best loss, and completed epoch count.

CheckpointDescription
latest.ptUpdated after every epoch
best.ptUpdated when the monitored validation loss improves
epoch_NNNN.ptSaved according to checkpoint_interval

Evaluate a saved or newly trained model against the configured validation dataset:

python
metrics = trainer.evaluate()
print(metrics)

Evaluation reports validation loss and COCO bounding-box metrics. Segmentation models additionally report mask metrics.

Resume Training ​

Recreate the original training setup, set epochs to the new total number of epochs, and resume from a compatible checkpoint.

Recreate the Model and Datasets ​

Construct the same model architecture and datasets used by the original run. Keep the model variant, resolution, class configuration, transforms, and mask requirements unchanged.

python
model = RFDETR(
    variant="seg-nano",
    num_classes=max(category_ids),
    pretrained=True,
)

train_dataset = COCODataset(
    "dataset/train",
    transforms=model.train_transforms,
    include_masks=model.requires_masks,
)
valid_dataset = COCODataset(
    "dataset/valid",
    transforms=model.val_transforms,
    include_masks=model.requires_masks,
)

Instantiate the Trainer ​

Set epochs to the total desired epoch count, including the epochs already completed by the checkpoint.

python
trainer = Trainer(
    model=model,
    training_dataset=train_dataset,
    validation_dataset=valid_dataset,
    output_dir="results/training",
    epochs=40,
    batch_size=4,
    learning_rate=1e-4,
)

Resume from Checkpoint ​

python
history = trainer.resume("results/training/latest.pt")

The trainer restores the saved state and continues from the next epoch until it reaches the configured total.

Checkpoint Compatibility

The model architecture, variant, resolution, and class configuration must remain compatible with the checkpoint being restored.

Trainer Reference ​

Constructor Options ​

OptionDefaultDescription
modelrequiredIris TrainableModel or raw PyTorch module
training_datasetrequiredDataset used for optimization
output_dirrequiredDirectory for checkpoints and training outputs
criterionNoneExplicit loss criterion; otherwise supplied by the adapter or model
adapterNoneModel-family TrainingAdapter for a raw PyTorch model
model_nameNonePublic model identifier stored in checkpoints for export
validation_datasetNoneOptional dataset used for validation and evaluation
optimizerNoneCustom optimizer; defaults to AdamW
schedulerNoneOptional learning-rate scheduler
epochs10Total epoch count, including restored epochs
batch_size4Samples per batch
learning_rate1e-4Learning rate for the default AdamW optimizer
weight_decay1e-4Weight decay for the default AdamW optimizer
deviceNoneExplicit device; automatically selects CUDA when available
num_workers0DataLoader worker processes
gradient_clip_norm0.1Maximum gradient norm; None disables clipping
mixed_precisionTrueEnables automatic mixed precision on CUDA
evaluate_masksNoneExplicitly enables mask training and metrics; otherwise inferred
evaluation_interval1Epochs between validation runs
checkpoint_interval1Epochs between numbered checkpoints
seed42Python, NumPy, and PyTorch random seed
metric_loggerNoneOptional TensorBoardMetricLogger
on_epoch_endNoneCallback invoked after every epoch

Methods ​

MethodReturnsDescription
Trainer.collate_fn(samples)BatchKeeps variable-sized image tensors and their target dictionaries as separate sequences.
trainer.train(resume_from=None)list[dict[str, float]]Runs training and validation, writes checkpoints, logs metrics, and returns one metric dictionary per completed epoch.
trainer.evaluate(dataset=None)dict[str, float]Evaluates the supplied dataset, or the configured validation dataset, without updating the model.
trainer.resume(checkpoint_path)list[dict[str, float]]Restores a checkpoint and continues training until the configured total number of epochs is reached.
trainer.export(output_dir=None, artifact_name="model", from_best=True)PathExports best.pt or latest.pt as the deployment artifact for the configured model family.
trainer.save_checkpoint(checkpoint_path, epoch=...)NoneSaves model, criterion, optimizer, scheduler, scaler, metric history, and training progress.
trainer.load_checkpoint(checkpoint_path)NoneRestores the saved training state and sets the next epoch from which training continues.

Next Steps ​