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.
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
from telekinesis.iris.models import RFDETR
model = RFDETR(
variant="seg-nano",
num_classes=max(category_ids),
pretrained=True,
)| Task | Available variants |
|---|---|
| Object detection | nano, small, medium, base, large |
| Instance segmentation | seg-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
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:
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.
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.
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
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.
| Checkpoint | Description |
|---|---|
latest.pt | Updated after every epoch |
best.pt | Updated when the monitored validation loss improves |
epoch_NNNN.pt | Saved according to checkpoint_interval |
Evaluate a saved or newly trained model against the configured validation dataset:
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.
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.
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
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
| Option | Default | Description |
|---|---|---|
model | required | Iris TrainableModel or raw PyTorch module |
training_dataset | required | Dataset used for optimization |
output_dir | required | Directory for checkpoints and training outputs |
criterion | None | Explicit loss criterion; otherwise supplied by the adapter or model |
adapter | None | Model-family TrainingAdapter for a raw PyTorch model |
model_name | None | Public model identifier stored in checkpoints for export |
validation_dataset | None | Optional dataset used for validation and evaluation |
optimizer | None | Custom optimizer; defaults to AdamW |
scheduler | None | Optional learning-rate scheduler |
epochs | 10 | Total epoch count, including restored epochs |
batch_size | 4 | Samples per batch |
learning_rate | 1e-4 | Learning rate for the default AdamW optimizer |
weight_decay | 1e-4 | Weight decay for the default AdamW optimizer |
device | None | Explicit device; automatically selects CUDA when available |
num_workers | 0 | DataLoader worker processes |
gradient_clip_norm | 0.1 | Maximum gradient norm; None disables clipping |
mixed_precision | True | Enables automatic mixed precision on CUDA |
evaluate_masks | None | Explicitly enables mask training and metrics; otherwise inferred |
evaluation_interval | 1 | Epochs between validation runs |
checkpoint_interval | 1 | Epochs between numbered checkpoints |
seed | 42 | Python, NumPy, and PyTorch random seed |
metric_logger | None | Optional TensorBoardMetricLogger |
on_epoch_end | None | Callback invoked after every epoch |
Methods
| Method | Returns | Description |
|---|---|---|
Trainer.collate_fn(samples) | Batch | Keeps 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) | Path | Exports best.pt or latest.pt as the deployment artifact for the configured model family. |
trainer.save_checkpoint(checkpoint_path, epoch=...) | None | Saves model, criterion, optimizer, scheduler, scaler, metric history, and training progress. |
trainer.load_checkpoint(checkpoint_path) | None | Restores the saved training state and sets the next epoch from which training continues. |