myria3d.callbacks

Submodules

myria3d.callbacks.comet_callbacks

class myria3d.callbacks.comet_callbacks.LogCode(code_dir: str)[source]

Bases: Callback

Upload all code files to comet, at the beginning of the run.

on_train_start(trainer, pl_module)[source]

Called when the train begins.

class myria3d.callbacks.comet_callbacks.LogLogsPath[source]

Bases: Callback

Logs run working directory to comet.ml

setup(trainer, pl_module, stage)[source]

Called when fit, validate, test, predict, or tune begins.

myria3d.callbacks.comet_callbacks.get_comet_logger(trainer: Trainer) → CometLogger | None[source]

Safely get logger from Trainer. If there is no comet logger, simply returns None to deactivate comet-based callbacks.

myria3d.callbacks.comet_callbacks.log_comet_cm(pl_module, confmat, phase, class_names)[source]

Method used in the metric logging callback.

myria3d.callbacks.finetuning_callbacks

class myria3d.callbacks.finetuning_callbacks.FinetuningFreezeUnfreeze(d_in: int = 9, num_classes: int = 6, unfreeze_fc_end_epoch: int = 3, unfreeze_decoder_train_epoch: int = 6)[source]

Bases: BaseFinetuning

finetune_function(pl_module, current_epoch, optimizer, optimizer_idx)[source]

Unfreeze layers sequentially, starting from the end of the architecture.

freeze_before_training(pl_module)[source]

Update in and out dimensions, and freeze everything at start.

myria3d.callbacks.metric_callbacks

class myria3d.callbacks.metric_callbacks.ModelMetrics(num_classes=7)[source]

Bases: Callback

Compute metrics for multiclass classification.

Accuracy, Precision, Recall, F1Score are micro-averaged. IoU (Jaccard Index) is macro-average to get the mIoU. All metrics are also computed per class.

Be careful when manually computing/reseting metrics. See: https://lightning.ai/docs/torchmetrics/stable/pages/lightning.html

on_test_batch_end(trainer, pl_module, outputs, batch, batch_idx)[source]

Called when the test batch ends.

on_test_epoch_end(trainer, pl_module)[source]

Called when the test epoch ends.

on_train_batch_end(trainer, pl_module, outputs, batch, batch_idx)[source]

Called when the train batch ends.

Note

The value outputs["loss"] here will be the normalized value w.r.t accumulate_grad_batches of the loss returned from training_step.

on_train_epoch_end(trainer, pl_module)[source]

Called when the train epoch ends.

To access all batch outputs at the end of the epoch, you can cache step outputs as an attribute of the pytorch_lightning.core.LightningModule and access them in this hook:

class MyLightningModule(L.LightningModule):
    def __init__(self):
        super().__init__()
        self.training_step_outputs = []

    def training_step(self):
        loss = ...
        self.training_step_outputs.append(loss)
        return loss


class MyCallback(L.Callback):
    def on_train_epoch_end(self, trainer, pl_module):
        # do something with all training_step outputs, for example:
        epoch_mean = torch.stack(pl_module.training_step_outputs).mean()
        pl_module.log("training_epoch_mean", epoch_mean)
        # free up the memory
        pl_module.training_step_outputs.clear()
on_val_epoch_end(trainer, pl_module)[source]
on_validation_batch_end(trainer, pl_module, outputs, batch, batch_idx)[source]

Called when the validation batch ends.

Module contents