# luxonis_lightning

Python API: `luxonis_train.lightning.luxonis_lightning`

The Lightning module that runs the node graph.

It builds the nodes, runs them in topological order, computes the losses and the metrics, and decides what to log on each epoch.

## Classes

### LuxonisLightningModule

The Lightning module that holds the whole model.

The module builds every node of the config and keeps them in `nodes`, together with the losses, the metrics, and the visualizers
attached to each node. The model topology is an acyclic graph of nodes. `nodes.graph` stores it as a mapping from a node name to
the names of the nodes that feed it.

[full_forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
runs the graph. Lightning drives training, validation, testing, and prediction through the `*_step` methods and the `on_*` hooks.
[LuxonisModel](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/core/core.md)
owns the module and attaches it to a trainer.

#### Methods

##### init

```python
def __init__(cfg: Config, save_dir: PathType, input_shapes: dict[str, Size], dataset_metadata: DatasetMetadata | None = None, *, _core: luxonis_train.core.LuxonisModel | None = None, **kwargs):
```

Build the module from a config.

The constructor does the following steps:

 * It builds the nodes and their attached modules.
 * It builds the training strategy.
 * It loads the checkpoint that `cfg.model.weights` names, when there is one. That load leaves the module in evaluation mode.
 * It saves the versions of `luxonis_train` and `luxonis_ml` as hyperparameters.

Parameters

 * `cfg` (`Config`): The config that defines the model, the loader, and the trainer.
 * `save_dir` (`PathType`): The directory where the checkpoints and the logs go, as a string or a `pathlib.Path`.
 * `input_shapes` (`dict[str, Size]`): The shape of every loader input, keyed by input name and without the batch dimension.
 * `dataset_metadata` (`DatasetMetadata | None`): The metadata of the dataset. `None` builds an empty
   [DatasetMetadata](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/utils/dataset_metadata.md).
 * `_core` (`luxonis_train.core.LuxonisModel | None`): The
   [LuxonisModel](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/core/core.md)
   that owns this module.
   [core](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
   raises without it.
 * `**kwargs`: Extra keyword arguments for the `LightningModule` constructor.

##### compute_training_loss

```python
def compute_training_loss(train_batch: tuple[dict[str, Tensor] | Tensor, Labels]) -> Tensor:
```

Run one training batch and return its total loss.

The method runs
[full_forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
with the losses on, sums them with
[compute_losses](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/utils.md),
and records every loss value in the training loss accumulator that
[on_train_epoch_end](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
logs.

Parameters

 * `train_batch` (`tuple[dict[str, Tensor] | Tensor, Labels]`): The loader inputs and the labels of the batch.

Returns

 * `Tensor`: The sum of every loss, each already scaled by its config `weight`, of shape `[1]` on the device of the module.

Raises

 * `ValueError`: When no node produced a loss.

##### configure_callbacks

```python
def configure_callbacks(self) -> list[pl.Callback]:
```

Build the callbacks that every run gets.

Lightning calls it at the start of every `fit`, `validate`, `test`, or `predict` call of the trainer. The method returns
[Nodes.build_callbacks](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/utils.md)
of `save_dir`, which holds, in this order:

 * a
   [TrainingManager](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/training_manager.md);
 * a
   [LuxonisModelSummary](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/luxonis_model_summary.md);
 * a `ModelCheckpoint` on `val/loss`, saved in `save_dir / "min_val_loss"`;
 * an
   [AIMETCallback](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/aimet_callback.md),
   when `cfg.exporter.aimet.active` is on;
 * a `ModelCheckpoint` on the main metric, saved in `save_dir / "best_val_metric"`, when the config has one;
 * every active callback of `cfg.trainer.callbacks`;
 * a `GradientAccumulationScheduler`, when `cfg.trainer.accumulate_grad_batches` is set and no callback of that class is present
   yet.

Returns

 * `list[pl.Callback]`: The callbacks, in that order.

##### configure_optimizers

```python
def configure_optimizers(self) -> tuple[Sequence[Optimizer], Sequence[LRSchedulerTypeUnion | LRSchedulerConfig]]:
```

Build the optimizers and the schedulers of the run.

Lightning calls it when a fit starts. The method builds a
[TrainingPlan](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/training_plan.md)
with
[resolve_training_plan](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/training_plan.md).
The plan combines the finetuning rules of the nodes, the rules of the training strategy, and the trainer-level optimizer and
scheduler. The method instantiates the plan with
[build_training_plan](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/training_plan.md).
It then hands the group handles to the training strategy and to the freeze schedule of the nodes. It stores the runtime in
[training_plan](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md),
and logs a summary of the parameter groups.

A `ReduceLROnPlateau` scheduler in `max` mode monitors the main metric as `val/metric/<task>-<node>/<metric>`. Without a main
metric,
[build_training_plan](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/training_plan.md)
raises `ValueError`. In any other mode the scheduler monitors `val/loss`.

Returns

 * `tuple[Sequence[Optimizer], Sequence[LRSchedulerTypeUnion | LRSchedulerConfig]]`: A list with the one optimizer of the plan,
   which is a
   [CompositeOptimizer](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/optimizers/composite_optimizer.md)
   when the plan has several inner optimizers, and the scheduler configs of the plan.

##### detach

```python
def detach(self):
```

Detach the module from its trainer.

The method sets `trainer` to `None`. After it,
[tracker](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
returns `None` and
[progress_bar](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
raises `AttributeError`.

##### export_onnx

```python
def export_onnx(save_path: PathType, **kwargs) -> Path:
```

Export the model to ONNX.

The method puts the module in evaluation mode, moves it to the CPU, and enters export mode. It builds a zero input of shape `[1,
*shape]` for every loader input that a node reads. It runs the graph once to name the outputs, and calls `to_onnx` of the
Lightning module. It then logs the path, leaves export mode, puts the module in training mode, and moves it back to its device.

The default `input_names` are the loader input names. The default `output_names` are the `export_output_names` of a node when
their count matches its outputs, and `<task>/<node>/<output>/<index>` otherwise. A count mismatch logs a warning. On PyTorch 2.5
and later, `dynamo` defaults to `False`.

Parameters

 * `save_path` (`PathType`): The path of the ONNX file, as a string or a `pathlib.Path`.
 * `**kwargs`: Extra keyword arguments for `torch.onnx.export`, such as `opset_version` and `dynamic_axes`.

Returns

 * `Path`: `save_path` as a `pathlib.Path`.

##### forward

```python
def forward(inputs: dict[str, Tensor] | Tensor) -> tuple[Tensor, ...]:
```

Run the graph and return the outputs as a flat tuple.

This is the entry point of a direct call of the module and of the ONNX export. It runs
[full_forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
without the losses, the metrics, and the visualizations, and it flattens the packets of the output nodes.

The tuple is sorted by node name, then by output key, then by index. A list value contributes each of its tensors. A single tensor
value repeats once for every entry of its first dimension. The export uses a batch of one.

Parameters

 * `inputs` (`dict[str, Tensor] | Tensor`): The loader inputs, keyed by input name. A bare tensor is the input named
   `image_source`.

Returns

 * `tuple[Tensor, ...]`: The flattened outputs of the output nodes.

##### full_forward

```python
def full_forward(inputs: dict[str, Tensor] | Tensor, labels: Labels | None = None, images: Tensor | None = None, *, compute_loss: bool = True, compute_metrics: bool = False, compute_visualizations: bool = False) -> LuxonisOutput:
```

Run every node of the graph and collect its results.

The nodes run in topological order. A node runs after every node that feeds it. The method skips a node whose `export` flag and
`remove_on_export` are both set. The method drops the output of a node from memory once no later node reads it, unless the node is
an output node.

After each node runs, its attached modules move to the device of the module and run:

 * With `compute_loss` on and `labels` given, every loss of the node runs on the outputs and the labels. While the module is in
   training mode, the method puts the loss in training mode first.
 * With `compute_metrics` on and `labels` given, every metric of the node updates its state. The method does not return the metric
   values.
 * With `compute_visualizations` on and `images` given, every visualizer of the node draws on `images`, and
   [combine_visualizations](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/attached_modules/visualizers/utils.md)
   merges its result into one image batch.

Parameters

 * `inputs` (`dict[str, Tensor] | Tensor`): The loader inputs, keyed by input name. A bare tensor is the input named
   `image_source`. The tensors move to the device of the module.
 * `labels` (`Labels | None`): The labels of the batch, keyed as `<task>/<label>`. The tensors move to the device of the module.
   `None` skips the losses and the metrics.
 * `images` (`Tensor | None`): The denormalized images of the batch, of shape `[B, C, H, W]`, that the visualizers draw on. `None`
   skips the visualizations.
 * `compute_loss` (`bool`): Whether to run the losses.
 * `compute_metrics` (`bool`): Whether to update the metrics.
 * `compute_visualizations` (`bool`): Whether to run the visualizers.

Returns

 * `LuxonisOutput`: `outputs` holds the packet of every output node that ran, keyed by node name. `losses` holds the value of
   every loss, keyed by node name and loss name; a value is a tensor, or a tuple of the tensor and its sub-losses.
   `visualizations` holds the image batch of every visualizer, of shape `[B, C, H, W]`, keyed by node name and visualizer name.
   `metrics` stays empty.

##### get_mlflow_logging_keys

```python
def get_mlflow_logging_keys(self) -> dict[str, list[str]]:
```

Return the metric keys and the artifact paths of a full run.

The result predicts what the tracker receives over a run with the current config.
[LuxonisModel.tune](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/core/core.md)
reads it to find the monitor key of the main metric. The method calls `compute` on every metric to learn the names of its
sub-metrics, without a reset.

In every key, `<task>-<node>` is the log name of the node, see
[Nodes.formatted_name](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/utils.md).

The `"metrics"` list holds:

 * `<mode>/loss` and `<mode>/loss/<task>-<node>/<loss>` for `train`, `val`, and `test`. The sub-loss keys are not listed;
 * `val/metric/<task>-<node>/<name>` and `test/metric/<task>-<node>/<name>` for every metric value whose name does not contain
   `confusion_matrix`. The `val` key is absent when the run has no validation epoch, that is, when `validation_interval` is `-1`
   and `run_validation_after_first_epoch` is off;
 * the timing keys of
   [TrainingProgressCallback](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/training_progress_callback.md),
   when the config lists it.

The `"artifacts"` list holds:

 * `<mode>/metrics/<epoch>/<task>-<node>/<name>.json` for every metric value whose name contains `confusion_matrix`, at epoch 0
   and at every validation epoch for `val`, and at epoch `cfg.trainer.epochs` for `test`;
 * `<mode>/metrics/<task>-<node>/<metric>/<epoch>/<artifact>.png` for every artifact name of every metric, at the same epochs;
 * `<mode>/visualizations/<task>-<node>/<visualizer>/<epoch>/<i>.png` for every visualizer, at the same epochs, for `<i>` below
   `cfg.trainer.n_log_images`;
 * the files of the callbacks the config lists, active or not. `<name>` is `cfg.exporter.name`, or `cfg.model.name` without it: *
   `UploadCheckpoint`: `best_val_metric.ckpt` and `min_val_loss.ckpt`;
    * `ExportOnTrainEnd`: `<name>.onnx`;
    * `ArchiveOnTrainEnd`: `<name>.onnx.tar.xz`;
    * `ConvertOnTrainEnd`: `<name>.onnx` and `<name>.onnx.tar.xz`;
    * `AIMETCallback`: `<name>.onnx`, `<name>.onnx.data`, `<name>.onnx.tar.xz`, and `<name>.encodings`;
 * `luxonis_train.log`, `training_config.yaml`, and `<model name>.yaml`, where `<model name>` is `cfg.model.name`.

Returns

 * `dict[str, list[str]]`: The `"metrics"` and the `"artifacts"` lists, each sorted.

##### load_checkpoint

```python
def load_checkpoint(ckpt: PathType | dict[str, Any] | None):
```

Load the weights of every node from a checkpoint.

The method loads the checkpoint file on the device of the module when it gets a path. It warns when `cfg.trainer.resume_training`
is on and the `trainer.epochs` of the checkpoint config is above `cfg.trainer.epochs`. It also warns when the predefined model of
the config and the one of the checkpoint resolve to different classes. It warns too when the predefined model of the checkpoint no
longer resolves.

The method splits the state dict by node. A checkpoint of version 0.4 or later keys a node as `nodes.<node>.module.`, an older one
as `nodes.<node>.`. A checkpoint without a `version` key counts as `0.3.0`. Each part loads into its node with `strict=True`. When
that fails:

 * With `cfg.trainer.strict_weights_loading` on, the method raises the error.
 * Otherwise, the method remaps the keys through the `execution_order` of the checkpoint and of the model, and loads the part
   again. Without an execution order, or when the remap fails, the part loads with `strict=False` and a log message reports it.

The method runs a forward pass on zero inputs to record the execution order of the model, which leaves the module in evaluation
mode.

Parameters

 * `ckpt` (`PathType | dict[str, Any] | None`): A path to a checkpoint file, or a loaded checkpoint dictionary with a `state_dict`
   key. `None` loads nothing.

Raises

 * `ValueError`: When the checkpoint has no `state_dict` key.
 * `RuntimeError`: When a node fails to load with `cfg.trainer.strict_weights_loading` on, or when the checkpoint holds no key for
   a node.

##### load_state_dict

```python
def load_state_dict(state_dict: Mapping[str, Tensor], strict: bool = True) -> _IncompatibleKeys:
```

Load a state dict, and relax the check when a run resumes.

Lightning calls this method when it restores a run from a checkpoint, with `strict` set to the `strict_loading` of the module,
`True` by default. When `cfg.trainer.resume_training` is off, the method behaves like `torch.nn.Module.load_state_dict`.

When `cfg.trainer.resume_training` is on:

 * With `cfg.trainer.strict_weights_loading` off, the method loads the state dict with `strict=False` and returns every mismatch.
 * With `cfg.trainer.strict_weights_loading` on, the method leaves out the keys that point from a loss, a metric, or a visualizer
   back at its node. Such a key starts with `nodes.<node>.losses.`, `nodes.<node>.metrics.`, or `nodes.<node>.visualizers.` and
   holds `_node.`. The method ignores the same keys in the result. It raises for any other missing or unexpected key.

Parameters

 * `state_dict` (`Mapping[str, Tensor]`): The parameters and the buffers to load, keyed by their name in this module.
 * `strict` (`bool`): Whether every key must match. Ignored when a run resumes.

Returns

 * `_IncompatibleKeys`: The `missing_keys` and the `unexpected_keys` of the load. Both are empty after a strict resume.

Raises

 * `RuntimeError`: When a run resumes with strict weight loading and a key that is not such an attached-module key is missing or
   unexpected.

##### on_save_checkpoint

```python
def on_save_checkpoint(checkpoint: dict[str, Any]):
```

Add the metadata of the run to a checkpoint.

Lightning calls it before it writes a checkpoint. The method changes `checkpoint` in place:

 * It drops the keys of `state_dict` that point from a loss, a metric, or a visualizer back at its node.
 * It adds `version`, the version of `luxonis_train`.
 * It adds `execution_order`, the names of the leaf modules that hold parameters, in the order a forward pass runs them.
 * It adds `config`, the dump of the config.
 * It adds `dataset_metadata`, the dump of the dataset metadata.
 * It adds `predefined_model`: the one of the config, with `latest` resolved to a version number, or else the pin inherited from
   the loaded checkpoint. Without either, the method removes the key.

The execution order comes from a forward pass on zero inputs, which leaves the module in evaluation mode.

Parameters

 * `checkpoint` (`dict[str, Any]`): The checkpoint dictionary that Lightning is about to write.

##### on_test_epoch_end

```python
def on_test_epoch_end(self):
```

Log the test results of the epoch.

Lightning calls it at the end of the test epoch. The method does the same as
[on_validation_epoch_end](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
with the `test` prefix, without the change of the validation interval.

Raises

 * `RuntimeError`: When the run uses DDP and a metric value is not on the device of the module.

##### on_train_epoch_end

```python
def on_train_epoch_end(self):
```

Log the mean training losses of the epoch.

Lightning calls it at the end of every training epoch. The method logs every entry of the training loss accumulator as
`train/<name>` with `sync_dist=True`. It replaces the node name in `<name>` with the log name of the node, see
[Nodes.formatted_name](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/utils.md).
It then clears the accumulator.

##### on_train_epoch_start

```python
def on_train_epoch_start(self):
```

Tell every node the number of the epoch that starts.

Lightning calls it at the start of every training epoch. The method sets `current_epoch` on the module of every node, which the
node and its attached modules read.

##### on_validation_epoch_end

```python
def on_validation_epoch_end(self):
```

Log the validation results of the epoch.

Lightning calls it at the end of every validation epoch, the sanity check included. The method does the following steps:

 * It logs the mean validation losses as `val/<name>`, as
   [on_train_epoch_end](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
   does for training.
 * It computes every metric. On the main process, outside the sanity check, it uploads the image artifacts of the metric to the
   tracker. It then resets the metric. It logs a scalar value as `val/metric/<task>-<node>/<name>` with `sync_dist=True`. A 2-D
   value goes to the tracker as a matrix named `val/metrics/<epoch>/<task>-<node>/<name>`, with the row labels of
   [BaseLuxonisProgressBar.format_matrix_for_printing](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/luxonis_progress_bar.md)
   as its class names. `<task>-<node>` is the log name of the node, see
   [Nodes.formatted_name](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/utils.md).
   `<name>` is the metric identifier, or the name of a sub-metric when `cfg.trainer.log_sub_metrics` is on.
 * On the main process, it logs the loss, prints the metrics through the progress bar, and logs the main metric value.
 * When the epoch logged fewer than `cfg.trainer.n_log_images` images, it logs a warning. It then logs the visualizations that the
   epoch buffered. When the count matches, it stops the buffering of skipped images for the rest of the run.
 * It clears the buffered visualizations, the image counters, and the loss accumulator.

After the validation of epoch 0, outside the sanity check, it puts back the `trainer.check_val_every_n_epoch` that
[setup](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
replaced.

Raises

 * `RuntimeError`: When the run uses DDP and a metric value is not on the device of the module.

##### predict_step

```python
def predict_step(batch: tuple[dict[str, Tensor] | Tensor, Labels]) -> LuxonisOutput:
```

Run one batch of the prediction loop.

Lightning calls it once for every batch of the prediction loader. The method denormalizes the image input into `uint8` images. It
then runs
[full_forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
on the batch with those images, the visualizers on, and the losses and the metrics off.

Parameters

 * `batch` (`tuple[dict[str, Tensor] | Tensor, Labels]`): The loader inputs and the labels of the batch.

Returns

 * `LuxonisOutput`: The packet of every output node and the image of every visualizer. `losses` and `metrics` are empty.

##### reparameterize

```python
def reparameterize(self) -> Self:
```

Reparameterize every
[Reparameterizable](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/blocks/reparameterizable.md)
block of the model.

Unlike
[set_export_mode](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md),
the method leaves the export mode of the nodes unchanged.
[set_export_mode](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
with `False` restores the blocks.
[LuxonisModel.quantize](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/core/core.md)
calls it before the quantization.

Returns

 * `Self`: The module.

##### set_export_mode

```python
def set_export_mode(mode: bool) -> Self:
```

Switch every node into or out of export mode.

The method calls
[BaseNode.set_export_mode](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md)
on every
[BaseNode](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/base_node.md)
among the submodules of the module, nested ones included. A node reads its export flag to change its outputs, and the call
reparameterizes its
[Reparameterizable](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/nodes/blocks/reparameterizable.md)
blocks. With `False`, the call restores those blocks.

Parameters

 * `mode` (`bool`): `True` to enter export mode, `False` to leave it.

Returns

 * `Self`: The module.

##### setup

```python
def setup(stage: str):
```

Make the first epoch validate when the config asks for it.

Lightning calls it at the start of every stage. The method acts only when all of these hold:

 * The stage is `fit`.
 * `cfg.trainer.run_validation_after_first_epoch` is on.
 * The current epoch is 0.
 * `trainer.check_val_every_n_epoch` is set and above 1.

It acts once per run.

Lightning decides at the end of an epoch whether validation runs from `trainer.check_val_every_n_epoch`. The method stores that
value and sets it to `1`, so the first epoch validates even when `validation_interval` would skip it.
[on_validation_epoch_end](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
restores the stored value after that first validation, so the rest of the run follows the configured interval.

Parameters

 * `stage` (`str`): The stage that starts: `"fit"`, `"validate"`, `"test"`, or `"predict"`.

##### test_step

```python
def test_step(test_batch: tuple[dict[str, Tensor] | Tensor, Labels]) -> dict[str, Tensor]:
```

Run one batch of the test loop.

Lightning calls it once for every batch of the test loader. The method does the same as
[validation_step](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
with the `test` loss accumulator and the `test` image prefix.

Parameters

 * `test_batch` (`tuple[dict[str, Tensor] | Tensor, Labels]`): The loader inputs and the labels of the batch.

Returns

 * `dict[str, Tensor]`: The loss values of the batch on the CPU, keyed as in
   [validation_step](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md).

##### train

```python
def train(mode: bool = True) -> Self:
```

Set the training mode of the module and of every node.

Besides the recursion of `torch.nn.Module.train`, the method calls `train` on every
[NodeWrapper](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/utils.md).
The wrapper passes the mode to the node and to the losses, the metrics, and the visualizers attached to it.

Parameters

 * `mode` (`bool`): `True` for training mode, `False` for evaluation mode.

Returns

 * `Self`: The module.

##### training_step

```python
def training_step(train_batch: tuple[dict[str, Tensor] | Tensor, Labels]) -> Tensor:
```

Run one batch of the training loop.

Lightning calls it once for every batch of the training loader and backpropagates the returned loss.

Parameters

 * `train_batch` (`tuple[dict[str, Tensor] | Tensor, Labels]`): The loader inputs and the labels of the batch.

Returns

 * `Tensor`: The total loss from
   [compute_training_loss](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md).

##### validation_step

```python
def validation_step(val_batch: tuple[dict[str, Tensor] | Tensor, Labels]) -> dict[str, Tensor]:
```

Run one batch of the validation loop.

Lightning calls it once for every batch of the validation loader. The method runs
[full_forward](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
with the losses, the metrics, and the visualizers on. It records the loss values in the validation loss accumulator. It runs the
visualizers only while the epoch has logged fewer images than `cfg.trainer.n_log_images`. It logs their images to the tracker.
When a label key contains `/classification`, it balances the logged images across the classes and buffers the skipped images for
[on_validation_epoch_end](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md).
Otherwise it logs the first images of the epoch.

Parameters

 * `val_batch` (`tuple[dict[str, Tensor] | Tensor, Labels]`): The loader inputs and the labels of the batch.

Returns

 * `dict[str, Tensor]`: The loss values of the batch on the CPU, keyed `"loss"` for the total, `"loss/<node>/<loss>"` for each
   loss, and `"loss/<node>/<loss>/<sub>"` for each sub-loss when `cfg.trainer.log_sub_losses` is on.

#### Attributes

##### cfg

The config the module was built from.

##### core

The
[LuxonisModel](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/core/core.md)
that owns this module.

Raises

 * `ValueError`: When the module was built without `_core`.

##### dataset_metadata

The metadata of the dataset. An empty
[DatasetMetadata](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/utils/dataset_metadata.md)
when the constructor got none.

##### image_source

The name of the loader input that holds the image, from `cfg.loader.image_source`.

##### logger

##### nodes

The node wrappers, keyed by node name.

##### outputs

The names of the output nodes, from `cfg.model.outputs`.

##### progress_bar

The progress bar callback of the attached trainer.

[LuxonisModel](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/core/core.md)
registers a
[LuxonisRichProgressBar](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/luxonis_progress_bar.md)
or a
[LuxonisTQDMProgressBar](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/luxonis_progress_bar.md),
and the module prints the results of an evaluation epoch through it. Without an attached trainer, the lookup raises
`AttributeError`.

##### save_dir

The directory where the checkpoints and the logs go.

##### tracker

The logger of the attached trainer, as a
[LuxonisTrackerPL](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/utils/tracker.md).

[LuxonisModel](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/core/core.md)
creates the trainer with a
[LuxonisTrackerPL](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/utils/tracker.md)
as its logger. The value is `None` while no trainer is attached.

##### trainer

##### training_plan

The runtime that
[configure_optimizers](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
built, or `None`.

It holds the optimizers and the schedulers of the training plan. The value is `None` before the first
[configure_optimizers](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
call.

##### training_strategy

The strategy that `cfg.trainer.training_strategy` names, or `None` when the config names none. A strategy that predates the
rule-based API comes wrapped in a
[LegacyStrategyAdapter](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/strategies/legacy.md).
