# freezing

Python API: `luxonis_train.lightning.freezing`

The freeze schedule of the nodes, derived from the config.

The config alone defines the schedule. A node with `freezing.active` stays frozen for each epoch before its unfreeze epoch, and
trains from that epoch on.
[resolve_unfreeze_epoch](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/freezing.md)
computes the unfreeze epoch.

[FreezeSchedule.apply](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/freezing.md)
derives the state of the nodes from the epoch number, so a repeated call for the same epoch gives the same state.
[TrainingManager](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/training_manager.md)
calls it when a fit starts and at the start of each training epoch. A resumed run thus gets the correct state, and the checkpoint
needs no state of the callback:

 * A frozen parameter stays in its parameter group, because the training plan puts each parameter in one fixed group. The
   optimizer state thus loads as a plain `state_dict`.
 * Each epoch start sets `requires_grad` and the batch normalization state again from the schedule.
 * The optimizer and scheduler state dicts hold the learning rates. Lightning saves and restores them.

[FreezeSchedule.apply](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/freezing.md)
sets `lr_after_unfreeze` only on the unfreeze epoch itself. A run that resumes after that epoch does not set it again. The
checkpoint already holds the learning rate that the scheduler reached. A second update of the rate would break a scheduler that
computes each rate from the previous rate, such as `StepLR`.

## Classes

### FreezeSchedule

The freeze plans of all nodes with `freezing.active`.

[Nodes](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/utils.md)
builds the schedule after it builds the nodes, and stores it in `freeze_schedule`.
[TrainingManager](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/callbacks/training_manager.md)
applies it.

#### Methods

##### init

```python
def __init__(plans: list[NodeFreezePlan]):
```

Initialize the schedule.

Parameters

 * `plans` (`list[NodeFreezePlan]`): One plan for each node with a freeze schedule. The schedule keeps the list itself, not a
   copy.

##### apply

```python
def apply(epoch: int, runtime: TrainingPlanRuntime | None = None):
```

Set the state of each scheduled node for an epoch.

For a node that is frozen in `epoch`, the method sets `requires_grad` of each parameter and `track_running_stats` of each batch
normalization layer to `False`. For any other node, it restores the original values of the plan. The method logs an info message
for a node when a parameter changes from trainable to frozen, and another when a parameter changes back. A repeated call for the
same epoch gives the same state.

With `runtime`, the method also sets `lr_after_unfreeze` as the base learning rate of each group of a node, but only when `epoch`
is the unfreeze epoch of the node. The module docstring tells why a later epoch does not set the rate. The groups come from
[attach_group_handles](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/freezing.md).

> **Example**
> The example turns the logger off, so the info messages do not show:

```pycon
>>> from loguru import logger
>>> from torch import nn
>>> logger.disable("luxonis_train")
>>> module = nn.Sequential(nn.Linear(2, 2), nn.BatchNorm1d(2))
>>> plan = NodeFreezePlan.from_module(
...     "head", module, unfreeze_epoch=2, lr_after_unfreeze=None
... )
>>> schedule = FreezeSchedule([plan])
>>> schedule.apply(epoch=0)
>>> module[0].weight.requires_grad, module[1].track_running_stats
(False, False)
>>> schedule.apply(epoch=2)
>>> module[0].weight.requires_grad, module[1].track_running_stats
(True, True)
>>> logger.enable("luxonis_train")
```

Parameters

 * `epoch` (`int`): The epoch number, from `0`.
 * `runtime` (`TrainingPlanRuntime | None`): The optimizers and schedulers of the run. `None` changes no learning rate.

##### attach_group_handles

```python
def attach_group_handles(runtime: TrainingPlanRuntime):
```

Store in each plan the parameter groups of its node.

[LuxonisLightningModule.configure_optimizers](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/luxonis_lightning.md)
calls the method after it builds the optimizers. For each plan, the method takes the handles of the groups that hold parameters of
the node. It checks that each of these groups holds parameters of that node only. `lr_after_unfreeze` thus changes no learning
rate of another node.

Parameters

 * `runtime` (`TrainingPlanRuntime`): The optimizers and schedulers of the run.

Raises

 * `RuntimeError`: When a group of a scheduled node also holds parameters of another node.

##### from_nodes

```python
def from_nodes(nodes: Nodes) -> Self:
```

Build the schedule from the nodes of a model.

The method creates a plan with
[NodeFreezePlan.from_module](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/freezing.md)
for each node whose `unfreeze_after` is not `None`. Build the schedule before any freeze, so that the plans record the original
state.

Parameters

 * `nodes` (`Nodes`): The nodes of the model. The method reads the `name`, `module`, `unfreeze_after`, and `lr_after_unfreeze` of
   each
   [NodeWrapper](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/utils.md).

Returns

 * `Self`: The schedule, with the plans in the order of `nodes`.

##### is_frozen

```python
def is_frozen(node_name: str, epoch: int) -> bool:
```

Return whether a node is frozen in an epoch.

Parameters

 * `node_name` (`str`): The name of the node in the graph.
 * `epoch` (`int`): The epoch number, from `0`.

Returns

 * `bool`: `True` when the node has a plan and `epoch` comes before its unfreeze epoch. `False` for a node without a plan.

#### Attributes

##### plans

The plans of the schedule, one for each scheduled node.

The property returns the stored list, not a copy.

### NodeFreezePlan

The freeze schedule of one node, with its original state.

#### Methods

##### from_module

```python
def from_module(node_name: str, module: nn.Module, unfreeze_epoch: int, lr_after_unfreeze: float | None) -> Self:
```

Create a plan that records the original state of a module.

The plan stores the `requires_grad` value of each parameter and the `track_running_stats` value of each batch normalization layer.
An unfreeze restores these values. It does not make every parameter trainable. A parameter that the node itself freezes, or a
layer built with `track_running_stats=False`, thus keeps its setting. Call this method before any freeze.

> **Example**
> ```pycon
>>> from torch import nn
>>> module = nn.Linear(2, 2)
>>> _ = module.bias.requires_grad_(False)
>>> plan = NodeFreezePlan.from_module(
...     "head", module, unfreeze_epoch=3, lr_after_unfreeze=None
... )
>>> plan.original_requires_grad
[True, False]
>>> plan.is_frozen(2), plan.unfreezes_at(3)
(True, True)
```

Parameters

 * `node_name` (`str`): The name of the node in the graph.
 * `module` (`nn.Module`): The node.
 * `unfreeze_epoch` (`int`): The first epoch in which the node trains.
 * `lr_after_unfreeze` (`float | None`): The base learning rate of the node from the unfreeze epoch on. `None` keeps the rate that the scheduler reached.

Returns

 * `Self`: The plan, without group handles.

##### is_frozen

```python
def is_frozen(epoch: int) -> bool:
```

Return whether the node is frozen in an epoch.

Parameters

 * `epoch` (`int`): The epoch number, from `0`.

Returns

 * `bool`: `True` when `epoch` comes before `unfreeze_epoch`.

##### unfreezes_at

```python
def unfreezes_at(epoch: int) -> bool:
```

Return whether the node unfreezes in an epoch.

Parameters

 * `epoch` (`int`): The epoch number, from `0`.

Returns

 * `bool`: `True` when `epoch` is `unfreeze_epoch`.

#### Attributes

##### batch_norms

The batch normalization layers of the node.

##### group_handles

The parameter groups that hold the parameters of the node. [FreezeSchedule.attach_group_handles](https://docs.luxonis.com/software-v3/ai-inference/model-source/training/luxonis-train/luxonis-train-api-reference/lightning/freezing.md) sets them. The tuple is empty before that call.

##### lr_after_unfreeze

The base learning rate of the parameter groups of the node from the unfreeze epoch on. `None` keeps the rate that the scheduler reached.

##### node_name

The name of the node in the graph.

##### original_requires_grad

The `requires_grad` value of each parameter before any freeze, in the order of `parameters`.

##### original_track_running_stats

The `track_running_stats` value of each layer before any freeze, in the order of `batch_norms`.

##### parameters

The parameters of the node.

##### unfreeze_epoch

The first epoch in which the node trains.

## Functions

### resolve_unfreeze_epoch

```python
def resolve_unfreeze_epoch(freezing: FreezingConfig, total_epochs: int) -> int | None:
```

Resolve `freezing.unfreeze_after` to an epoch number.

An integer is the epoch number itself. A float is a share of `total_epochs`, truncated to an integer. `None` gives `total_epochs`, so the node stays frozen for the whole run.

> **Example**
> ```pycon
>>> from luxonis_train.config.config import FreezingConfig
>>> resolve_unfreeze_epoch(
...     FreezingConfig(active=True, unfreeze_after=0.25), 10
... )
2
>>> resolve_unfreeze_epoch(FreezingConfig(active=True), 10)
10
>>> print(resolve_unfreeze_epoch(FreezingConfig(), 10))
None
```

Parameters

 * `freezing` (`FreezingConfig`): The `freezing` section of a node config.
 * `total_epochs` (`int`): The number of epochs of the run, from `trainer.epochs`.

Returns

 * `int | None`: The first epoch in which the node trains. `None` when `freezing.active` is `False`.
