mipcandy.presets.segmentation#

Module Contents#

Classes#

API#

class mipcandy.presets.segmentation.DeepSupervisionWrapper(loss: torch.nn.Module, *, weight_factors: Sequence[float] | None = None)[source]#

Bases: mipcandy.common.Loss

forward(outputs: Sequence[torch.Tensor], targets: Sequence[torch.Tensor]) tuple[torch.Tensor, dict[str, float]][source]#
class mipcandy.presets.segmentation.SegmentationTrainer(trainer_folder: str | os.PathLike[str], dataloader: torch.utils.data.DataLoader[tuple[torch.Tensor, torch.Tensor]], validation_dataloader: torch.utils.data.DataLoader[tuple[torch.Tensor, torch.Tensor]], *, recoverable: bool = True, profiler: bool = False, device: torch.device | str = 'cpu', console: rich.console.Console = Console())[source]#

Bases: mipcandy.training.Trainer

Initialization

num_classes: int = 1#
include_background: bool = True#
deep_supervision: bool = False#
deep_supervision_scales: Sequence[float] | None = None#
deep_supervision_weights: Sequence[float] | None = None#
_save_preview(x: torch.Tensor, title: str, quality: float, *, is_label: bool = False) None[source]#
apply_non_linearity(x: torch.Tensor, channel_dim: int) torch.Tensor[source]#
save_preview(image: torch.Tensor, label: torch.Tensor, output: torch.Tensor, *, quality: float = 0.75) None[source]#
build_ema(model: torch.nn.Module) torch.nn.Module[source]#
build_criterion() torch.nn.Module[source]#
build_optimizer(params: mipcandy.types.Params) torch.optim.Optimizer[source]#
build_scheduler(optimizer: torch.optim.Optimizer, num_epochs: int) torch.optim.lr_scheduler.LRScheduler[source]#
backward(images: torch.Tensor, labels: torch.Tensor, toolbox: mipcandy.training.TrainerToolbox) tuple[float, dict[str, float]][source]#
static prepare_deep_supervision_targets(labels: torch.Tensor, output_shapes: list[tuple[int, ...]]) list[torch.Tensor][source]#
class_percentages(ids: torch.Tensor) dict[int, float][source]#
static format_class_percentages(percentages: dict[int, float], prefix: str) dict[str, float][source]#
validate_case(idx: int, image: torch.Tensor, label: torch.Tensor, toolbox: mipcandy.training.TrainerToolbox) tuple[float, dict[str, float], torch.Tensor][source]#