From 190c316392b156fbc5c1c5dc42a665786b1e353b Mon Sep 17 00:00:00 2001 From: blanky Date: Wed, 25 Mar 2026 17:47:04 +0530 Subject: [PATCH 1/4] fix: enable ruff-format for jupyter files in pre-commit configuration --- .pre-commit-config.yaml | 14 +++----------- 1 file changed, 3 insertions(+), 11 deletions(-) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index b5ffdc823370..7a9cf5456ca2 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -27,15 +27,7 @@ repos: - repo: https://github.com/astral-sh/ruff-pre-commit rev: v0.11.13 hooks: - # Run the linter - id: ruff-check - types_or: [python, pyi] - # TODO: Enable when black and ruff format converge - # Run the formatter - # - id: ruff-format - # types_or: [python, pyi] - - - repo: https://github.com/psf/black - rev: 24.10.0 - hooks: - - id: black + types_or: [python, pyi, jupyter] + - id: ruff-format + types_or: [python, pyi, jupyter] From 5ac8cc8d12a18024c829f38a5869e0103fc6965a Mon Sep 17 00:00:00 2001 From: blanky Date: Wed, 25 Mar 2026 18:05:46 +0530 Subject: [PATCH 2/4] Refactor: Refactored codebase in ruff style. --- .pre-commit-config.yaml | 2 +- docker/test_image.py | 6 +- examples/cifar10_qat/utils.py | 5 +- examples/mnist/mnist_with_clearml_logger.py | 16 +- examples/mnist/mnist_with_neptune_logger.py | 16 +- examples/mnist/mnist_with_tensorboard.py | 28 +- .../mnist/mnist_with_tensorboard_logger.py | 26 +- .../mnist/mnist_with_tensorboard_on_tpu.py | 28 +- examples/mnist/mnist_with_visdom_logger.py | 26 +- examples/mnist/mnist_with_wandb_logger.py | 20 +- examples/notebooks/Cifar100_bench_amp.ipynb | 6 +- .../Cifar10_Ax_hyperparam_tuning.ipynb | 325 ++++++++-------- .../notebooks/CycleGAN_with_nvidia_apex.ipynb | 307 ++++++++------- .../CycleGAN_with_torch_cuda_amp.ipynb | 305 ++++++++------- .../EfficientNet_Cifar100_finetuning.ipynb | 359 +++++++++--------- examples/notebooks/FashionMNIST.ipynb | 153 ++++---- examples/notebooks/FastaiLRFinder_MNIST.ipynb | 8 +- .../HandlersTimeProfiler_MNIST.ipynb | 4 +- examples/notebooks/TextTransformers.ipynb | 47 ++- examples/notebooks/VAE.ipynb | 121 +++--- examples/siamese_network/siamese_network.py | 2 +- ignite/contrib/engines/common.py | 4 +- ignite/contrib/handlers/base_logger.py | 2 +- ignite/contrib/handlers/clearml_logger.py | 2 +- ignite/contrib/handlers/lr_finder.py | 2 +- ignite/contrib/handlers/mlflow_logger.py | 2 +- ignite/contrib/handlers/neptune_logger.py | 2 +- ignite/contrib/handlers/param_scheduler.py | 2 +- ignite/contrib/handlers/polyaxon_logger.py | 2 +- ignite/contrib/handlers/tensorboard_logger.py | 2 +- ignite/contrib/handlers/time_profilers.py | 2 +- ignite/contrib/handlers/tqdm_logger.py | 2 +- ignite/contrib/handlers/visdom_logger.py | 2 +- ignite/contrib/handlers/wandb_logger.py | 2 +- ignite/contrib/metrics/average_precision.py | 2 +- ignite/contrib/metrics/cohen_kappa.py | 2 +- ignite/contrib/metrics/gpu_info.py | 2 +- .../contrib/metrics/precision_recall_curve.py | 2 +- .../metrics/regression/canberra_metric.py | 2 +- .../regression/fractional_absolute_error.py | 2 +- .../metrics/regression/fractional_bias.py | 2 +- .../geometric_mean_absolute_error.py | 2 +- .../geometric_mean_relative_absolute_error.py | 2 +- .../metrics/regression/manhattan_distance.py | 2 +- .../regression/maximum_absolute_error.py | 2 +- .../mean_absolute_relative_error.py | 2 +- .../contrib/metrics/regression/mean_error.py | 2 +- .../regression/mean_normalized_bias.py | 2 +- .../regression/median_absolute_error.py | 2 +- .../median_absolute_percentage_error.py | 2 +- .../median_relative_absolute_error.py | 2 +- ignite/contrib/metrics/regression/r2_score.py | 2 +- .../regression/wave_hedges_distance.py | 2 +- ignite/contrib/metrics/roc_auc.py | 2 +- ignite/distributed/comp_models/__init__.py | 6 +- ignite/engine/deterministic.py | 3 +- ignite/engine/engine.py | 6 +- ignite/handlers/ema_handler.py | 5 +- ignite/handlers/mlflow_logger.py | 3 +- ignite/handlers/neptune_logger.py | 3 +- ignite/handlers/param_scheduler.py | 5 +- ignite/handlers/time_profilers.py | 10 +- ignite/handlers/tqdm_logger.py | 4 +- ignite/metrics/nlp/bleu.py | 3 +- ignite/metrics/running_average.py | 3 +- ignite/utils.py | 2 +- tests/ignite/conftest.py | 4 +- tests/ignite/contrib/engines/test_common.py | 6 +- .../distributed/comp_models/test_native.py | 74 ++-- tests/ignite/distributed/test_auto.py | 12 +- tests/ignite/engine/test_deterministic.py | 48 +-- tests/ignite/engine/test_memory_leaks.py | 1 - tests/ignite/handlers/test_neptune_logger.py | 14 +- tests/ignite/handlers/test_param_scheduler.py | 12 +- tests/ignite/metrics/nlp/__init__.py | 6 +- tests/ignite/metrics/test_accumulation.py | 12 +- tests/ignite/metrics/test_accuracy.py | 48 +-- .../metrics/test_classification_report.py | 6 +- tests/ignite/metrics/test_confusion_matrix.py | 12 +- tests/ignite/metrics/test_loss.py | 12 +- .../metrics/test_mean_average_precision.py | 1 - .../test_multilabel_confusion_matrix.py | 12 +- tests/ignite/metrics/test_precision.py | 24 +- tests/ignite/metrics/test_recall.py | 25 +- tests/ignite/metrics/test_running_average.py | 18 +- .../test_top_k_categorical_accuracy.py | 12 +- 86 files changed, 1204 insertions(+), 1086 deletions(-) diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 7a9cf5456ca2..35fd24bb75a7 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -28,6 +28,6 @@ repos: rev: v0.11.13 hooks: - id: ruff-check - types_or: [python, pyi, jupyter] + types_or: [python, pyi] - id: ruff-format types_or: [python, pyi, jupyter] diff --git a/docker/test_image.py b/docker/test_image.py index 77e1790d9c4d..cb800c7f3a52 100644 --- a/docker/test_image.py +++ b/docker/test_image.py @@ -24,9 +24,9 @@ def check_package(package_name, expected_version=None): old_version = version version = version.split("+")[0] print(f"Transformed version: {old_version} -> {version}") - assert ( - version == expected_version - ), f"Version mismatch for package {package_name}: got {version} but expected {expected_version}" + assert version == expected_version, ( + f"Version mismatch for package {package_name}: got {version} but expected {expected_version}" + ) if __name__ == "__main__": diff --git a/examples/cifar10_qat/utils.py b/examples/cifar10_qat/utils.py index 2a16d28519ba..be2b472d5d7c 100644 --- a/examples/cifar10_qat/utils.py +++ b/examples/cifar10_qat/utils.py @@ -214,8 +214,9 @@ def __init__( replace_stride_with_dilation = [False, False, False] if len(replace_stride_with_dilation) != 3: raise ValueError( - "replace_stride_with_dilation should be None " - "or a 3-element tuple, got {}".format(replace_stride_with_dilation) + "replace_stride_with_dilation should be None or a 3-element tuple, got {}".format( + replace_stride_with_dilation + ) ) self.groups = groups self.base_width = width_per_group diff --git a/examples/mnist/mnist_with_clearml_logger.py b/examples/mnist/mnist_with_clearml_logger.py index 1faeb375d935..ff1f39f96925 100644 --- a/examples/mnist/mnist_with_clearml_logger.py +++ b/examples/mnist/mnist_with_clearml_logger.py @@ -1,15 +1,15 @@ """ - MNIST example with training and validation monitoring using ClearML. +MNIST example with training and validation monitoring using ClearML. - Requirements: - ClearML: `pip install clearml` +Requirements: + ClearML: `pip install clearml` - Usage: +Usage: - Run the example: - ```bash - python mnist_with_clearml_logger.py - ``` + Run the example: + ```bash + python mnist_with_clearml_logger.py + ``` """ from argparse import ArgumentParser diff --git a/examples/mnist/mnist_with_neptune_logger.py b/examples/mnist/mnist_with_neptune_logger.py index 244b24419a47..b1983de272b7 100644 --- a/examples/mnist/mnist_with_neptune_logger.py +++ b/examples/mnist/mnist_with_neptune_logger.py @@ -1,15 +1,15 @@ """ - MNIST example with training and validation monitoring using Neptune. +MNIST example with training and validation monitoring using Neptune. - Requirements: - Neptune: `pip install neptune` +Requirements: + Neptune: `pip install neptune` - Usage: +Usage: - Run the example: - ```bash - python mnist_with_neptune_logger.py - ``` + Run the example: + ```bash + python mnist_with_neptune_logger.py + ``` """ diff --git a/examples/mnist/mnist_with_tensorboard.py b/examples/mnist/mnist_with_tensorboard.py index 75658a50e128..f58f17199db7 100644 --- a/examples/mnist/mnist_with_tensorboard.py +++ b/examples/mnist/mnist_with_tensorboard.py @@ -1,18 +1,18 @@ """ - MNIST example with training and validation monitoring using Tensorboard. - Requirements: - TensorboardX (https://github.com/lanpa/tensorboard-pytorch): `pip install tensorboardX` - or PyTorch >= 1.2 which supports Tensorboard - Tensorboard: `pip install tensorflow` (or just install tensorboard without the rest of tensorflow) - Usage: - Start tensorboard: - ```bash - tensorboard --logdir=/tmp/tensorboard_logs/ - ``` - Run the example: - ```bash - python mnist_with_tensorboard.py --log_dir=/tmp/tensorboard_logs - ``` +MNIST example with training and validation monitoring using Tensorboard. +Requirements: + TensorboardX (https://github.com/lanpa/tensorboard-pytorch): `pip install tensorboardX` + or PyTorch >= 1.2 which supports Tensorboard + Tensorboard: `pip install tensorflow` (or just install tensorboard without the rest of tensorflow) +Usage: + Start tensorboard: + ```bash + tensorboard --logdir=/tmp/tensorboard_logs/ + ``` + Run the example: + ```bash + python mnist_with_tensorboard.py --log_dir=/tmp/tensorboard_logs + ``` """ from argparse import ArgumentParser diff --git a/examples/mnist/mnist_with_tensorboard_logger.py b/examples/mnist/mnist_with_tensorboard_logger.py index 9c6b07c0ed08..6d0623047bc5 100644 --- a/examples/mnist/mnist_with_tensorboard_logger.py +++ b/examples/mnist/mnist_with_tensorboard_logger.py @@ -1,21 +1,21 @@ """ - MNIST example with training and validation monitoring using TensorboardX and Tensorboard. +MNIST example with training and validation monitoring using TensorboardX and Tensorboard. - Requirements: - Optionally TensorboardX (https://github.com/lanpa/tensorboard-pytorch): `pip install tensorboardX` - Tensorboard: `pip install tensorflow` (or just install tensorboard without the rest of tensorflow) +Requirements: + Optionally TensorboardX (https://github.com/lanpa/tensorboard-pytorch): `pip install tensorboardX` + Tensorboard: `pip install tensorflow` (or just install tensorboard without the rest of tensorflow) - Usage: +Usage: - Start tensorboard: - ```bash - tensorboard --logdir=/tmp/tensorboard_logs/ - ``` + Start tensorboard: + ```bash + tensorboard --logdir=/tmp/tensorboard_logs/ + ``` - Run the example: - ```bash - python mnist_with_tensorboard_logger.py --log_dir=/tmp/tensorboard_logs - ``` + Run the example: + ```bash + python mnist_with_tensorboard_logger.py --log_dir=/tmp/tensorboard_logs + ``` """ import sys diff --git a/examples/mnist/mnist_with_tensorboard_on_tpu.py b/examples/mnist/mnist_with_tensorboard_on_tpu.py index 2ec65e9605e9..80b13ebdcfb9 100644 --- a/examples/mnist/mnist_with_tensorboard_on_tpu.py +++ b/examples/mnist/mnist_with_tensorboard_on_tpu.py @@ -1,18 +1,18 @@ """ - MNIST example with training and validation monitoring using Tensorboard on TPU - Requirements: - - PyTorch >= 1.5 - - PyTorch XLA >= 1.5 - - Tensorboard: `pip install tensorflow` (or just install tensorboard without the rest of tensorflow) - Usage: - Start tensorboard: - ```bash - tensorboard --logdir=/tmp/tensorboard_logs/ - ``` - Run the example: - ```bash - python mnist_with_tensorboard_on_tpu.py --log_dir=/tmp/tensorboard_logs - ``` +MNIST example with training and validation monitoring using Tensorboard on TPU +Requirements: + - PyTorch >= 1.5 + - PyTorch XLA >= 1.5 + - Tensorboard: `pip install tensorflow` (or just install tensorboard without the rest of tensorflow) +Usage: + Start tensorboard: + ```bash + tensorboard --logdir=/tmp/tensorboard_logs/ + ``` + Run the example: + ```bash + python mnist_with_tensorboard_on_tpu.py --log_dir=/tmp/tensorboard_logs + ``` """ from argparse import ArgumentParser diff --git a/examples/mnist/mnist_with_visdom_logger.py b/examples/mnist/mnist_with_visdom_logger.py index 11f29f92849e..e14775bdd508 100644 --- a/examples/mnist/mnist_with_visdom_logger.py +++ b/examples/mnist/mnist_with_visdom_logger.py @@ -1,21 +1,21 @@ """ - MNIST example with training and validation monitoring using Visdom. +MNIST example with training and validation monitoring using Visdom. - Requirements: - Visdom (https://github.com/facebookresearch/visdom.git): - `pip install git+https://github.com/facebookresearch/visdom.git` +Requirements: + Visdom (https://github.com/facebookresearch/visdom.git): + `pip install git+https://github.com/facebookresearch/visdom.git` - Usage: +Usage: - Start visdom server: - ```bash - visdom -logging_level 30 - ``` + Start visdom server: + ```bash + visdom -logging_level 30 + ``` - Run the example: - ```bash - python mnist_with_visdom_logger.py - ``` + Run the example: + ```bash + python mnist_with_visdom_logger.py + ``` """ from argparse import ArgumentParser diff --git a/examples/mnist/mnist_with_wandb_logger.py b/examples/mnist/mnist_with_wandb_logger.py index 169f7c12a31e..03041e5ca752 100644 --- a/examples/mnist/mnist_with_wandb_logger.py +++ b/examples/mnist/mnist_with_wandb_logger.py @@ -1,19 +1,19 @@ """ - MNIST example with training and validation monitoring using Weights & Biases +MNIST example with training and validation monitoring using Weights & Biases - Requirements: - Weights & Biases: `pip install wandb` +Requirements: + Weights & Biases: `pip install wandb` - Usage: +Usage: - Make sure you are logged into Weights & Biases (use the `wandb` command). + Make sure you are logged into Weights & Biases (use the `wandb` command). - Run the example: - ```bash - python mnist_with_wandb_logger.py - ``` + Run the example: + ```bash + python mnist_with_wandb_logger.py + ``` - Go to https://wandb.com and explore your experiment. + Go to https://wandb.com and explore your experiment. """ from argparse import ArgumentParser diff --git a/examples/notebooks/Cifar100_bench_amp.ipynb b/examples/notebooks/Cifar100_bench_amp.ipynb index 8a128ebf32a4..f17faab5e51c 100644 --- a/examples/notebooks/Cifar100_bench_amp.ipynb +++ b/examples/notebooks/Cifar100_bench_amp.ipynb @@ -70,6 +70,7 @@ "import torch\n", "import torchvision\n", "import ignite\n", + "\n", "torch.__version__, torchvision.__version__, ignite.__version__" ] }, @@ -87,8 +88,8 @@ "outputs": [], "source": [ "!git clone https://github.com/pytorch/ignite.git /tmp/ignite\n", - "scriptspath=\"/tmp/ignite/examples/cifar100_amp_benchmark/\"\n", - "setup=f\"cd {scriptspath} && export PYTHONPATH=$PWD:$PYTHONPATH\"" + "scriptspath = \"/tmp/ignite/examples/cifar100_amp_benchmark/\"\n", + "setup = f\"cd {scriptspath} && export PYTHONPATH=$PWD:$PYTHONPATH\"" ] }, { @@ -105,6 +106,7 @@ "outputs": [], "source": [ "from torchvision.datasets.cifar import CIFAR100\n", + "\n", "CIFAR100(root=\"/tmp/cifar100/\", train=True, download=True)" ] }, diff --git a/examples/notebooks/Cifar10_Ax_hyperparam_tuning.ipynb b/examples/notebooks/Cifar10_Ax_hyperparam_tuning.ipynb index def2b9756d6e..f857fa0ec60b 100644 --- a/examples/notebooks/Cifar10_Ax_hyperparam_tuning.ipynb +++ b/examples/notebooks/Cifar10_Ax_hyperparam_tuning.ipynb @@ -51,6 +51,7 @@ "outputs": [], "source": [ "import sys\n", + "\n", "sys.path.insert(0, \"../../\")" ] }, @@ -101,59 +102,74 @@ " \"\"\"\n", " From : https://github.com/davidcpage/cifar10-fast/blob/master/bag_of_tricks.ipynb\n", "\n", - " Batch norm seems to work best with batch size of around 32. The reasons presumably have to do \n", - " with noise in the batch statistics and specifically a balance between a beneficial regularising effect \n", + " Batch norm seems to work best with batch size of around 32. The reasons presumably have to do\n", + " with noise in the batch statistics and specifically a balance between a beneficial regularising effect\n", " at intermediate batch sizes and an excess of noise at small batches.\n", - " \n", - " Our batches are of size 512 and we can't afford to reduce them without taking a serious hit on training times, \n", - " but we can apply batch norm separately to subsets of a training batch. This technique, known as 'ghost' batch \n", - " norm, is usually used in a distributed setting but is just as useful when using large batches on a single node. \n", + "\n", + " Our batches are of size 512 and we can't afford to reduce them without taking a serious hit on training times,\n", + " but we can apply batch norm separately to subsets of a training batch. This technique, known as 'ghost' batch\n", + " norm, is usually used in a distributed setting but is just as useful when using large batches on a single node.\n", " It isn't supported directly in PyTorch but we can roll our own easily enough.\n", " \"\"\"\n", + "\n", " def __init__(self, num_features, num_splits, eps=1e-05, momentum=0.1, weight=True, bias=True):\n", " super(GhostBatchNorm, self).__init__(num_features, eps=eps, momentum=momentum)\n", " self.weight.data.fill_(1.0)\n", " self.bias.data.fill_(0.0)\n", " self.weight.requires_grad = weight\n", - " self.bias.requires_grad = bias \n", + " self.bias.requires_grad = bias\n", " self.num_splits = num_splits\n", - " self.register_buffer('running_mean', torch.zeros(num_features*self.num_splits))\n", - " self.register_buffer('running_var', torch.ones(num_features*self.num_splits))\n", + " self.register_buffer(\"running_mean\", torch.zeros(num_features * self.num_splits))\n", + " self.register_buffer(\"running_var\", torch.ones(num_features * self.num_splits))\n", "\n", " def train(self, mode=True):\n", " if (self.training is True) and (mode is False):\n", - " self.running_mean = torch.mean(self.running_mean.view(self.num_splits, self.num_features), dim=0).repeat(self.num_splits)\n", - " self.running_var = torch.mean(self.running_var.view(self.num_splits, self.num_features), dim=0).repeat(self.num_splits)\n", + " self.running_mean = torch.mean(self.running_mean.view(self.num_splits, self.num_features), dim=0).repeat(\n", + " self.num_splits\n", + " )\n", + " self.running_var = torch.mean(self.running_var.view(self.num_splits, self.num_features), dim=0).repeat(\n", + " self.num_splits\n", + " )\n", " return super(GhostBatchNorm, self).train(mode)\n", - " \n", + "\n", " def forward(self, input):\n", " N, C, H, W = input.shape\n", " if self.training or not self.track_running_stats:\n", " return F.batch_norm(\n", - " input.view(-1, C*self.num_splits, H, W), self.running_mean, self.running_var, \n", - " self.weight.repeat(self.num_splits), self.bias.repeat(self.num_splits),\n", - " True, self.momentum, self.eps).view(N, C, H, W) \n", + " input.view(-1, C * self.num_splits, H, W),\n", + " self.running_mean,\n", + " self.running_var,\n", + " self.weight.repeat(self.num_splits),\n", + " self.bias.repeat(self.num_splits),\n", + " True,\n", + " self.momentum,\n", + " self.eps,\n", + " ).view(N, C, H, W)\n", " else:\n", " return F.batch_norm(\n", - " input, self.running_mean[:self.num_features], self.running_var[:self.num_features], \n", - " self.weight, self.bias, False, self.momentum, self.eps)\n", + " input,\n", + " self.running_mean[: self.num_features],\n", + " self.running_var[: self.num_features],\n", + " self.weight,\n", + " self.bias,\n", + " False,\n", + " self.momentum,\n", + " self.eps,\n", + " )\n", "\n", - " \n", - "class IdentityResidualBlock(nn.Module):\n", "\n", - " def __init__(self, num_channels, \n", - " conv_ksize=3, conv_pad=1,\n", - " gbn_num_splits=16):\n", + "class IdentityResidualBlock(nn.Module):\n", + " def __init__(self, num_channels, conv_ksize=3, conv_pad=1, gbn_num_splits=16):\n", " super(IdentityResidualBlock, self).__init__()\n", " self.res1 = nn.Sequential(\n", " Conv2d(num_channels, num_channels, kernel_size=conv_ksize, padding=conv_pad, stride=1, bias=False),\n", " GhostBatchNorm(num_channels, num_splits=gbn_num_splits, weight=False),\n", - " nn.CELU(alpha=0.3) \n", + " nn.CELU(alpha=0.3),\n", " )\n", " self.res2 = nn.Sequential(\n", " Conv2d(num_channels, num_channels, kernel_size=conv_ksize, padding=conv_pad, stride=1, bias=False),\n", " GhostBatchNorm(num_channels, num_splits=gbn_num_splits, weight=False),\n", - " nn.CELU(alpha=0.3) \n", + " nn.CELU(alpha=0.3),\n", " )\n", "\n", " def forward(self, x):\n", @@ -161,18 +177,17 @@ " x = self.res1(x)\n", " x = self.res2(x)\n", " return x + residual\n", - " \n", "\n", - "# We override conv2d to get proper padding for kernel size = 2 \n", + "\n", + "# We override conv2d to get proper padding for kernel size = 2\n", "class Conv2d(nn.Conv2d):\n", - " \n", " def __init__(self, *args, **kwargs):\n", " super(Conv2d, self).__init__(*args, **kwargs)\n", " if self.kernel_size == (2, 2):\n", " self.forward = self.ksize_2_forward\n", " self.ksize_2_padding = (0, self.padding[0], 0, self.padding[1])\n", " self.padding = (0, 0)\n", - " \n", + "\n", " def ksize_2_forward(self, x):\n", " x = F.pad(x, pad=self.ksize_2_padding)\n", " return super(Conv2d, self).forward(x)" @@ -185,17 +200,15 @@ "outputs": [], "source": [ "class FastResNet(nn.Module):\n", - " \n", - " def __init__(self, num_classes=10, \n", - " fmap_factor=64, conv_ksize=3, conv_pad=1, \n", - " gbn_num_splits=512 // 32, \n", - " classif_scale=0.0625):\n", + " def __init__(\n", + " self, num_classes=10, fmap_factor=64, conv_ksize=3, conv_pad=1, gbn_num_splits=512 // 32, classif_scale=0.0625\n", + " ):\n", " super(FastResNet, self).__init__()\n", - " \n", + "\n", " self.prep = nn.Sequential(\n", " Conv2d(3, fmap_factor, kernel_size=conv_ksize, padding=conv_pad, stride=1, bias=False),\n", " GhostBatchNorm(fmap_factor, num_splits=gbn_num_splits, weight=False),\n", - " nn.CELU(alpha=0.3)\n", + " nn.CELU(alpha=0.3),\n", " )\n", "\n", " self.layer1 = nn.Sequential(\n", @@ -203,34 +216,31 @@ " nn.MaxPool2d(kernel_size=2),\n", " GhostBatchNorm(fmap_factor * 2, num_splits=gbn_num_splits, weight=False),\n", " nn.CELU(alpha=0.3),\n", - " IdentityResidualBlock(fmap_factor * 2,\n", - " conv_ksize=conv_ksize, conv_pad=conv_pad, \n", - " gbn_num_splits=gbn_num_splits)\n", + " IdentityResidualBlock(\n", + " fmap_factor * 2, conv_ksize=conv_ksize, conv_pad=conv_pad, gbn_num_splits=gbn_num_splits\n", + " ),\n", " )\n", - " \n", + "\n", " self.layer2 = nn.Sequential(\n", " Conv2d(fmap_factor * 2, fmap_factor * 4, kernel_size=conv_ksize, padding=conv_pad, stride=1, bias=False),\n", " nn.MaxPool2d(kernel_size=2),\n", " GhostBatchNorm(fmap_factor * 4, num_splits=gbn_num_splits, weight=False),\n", - " nn.CELU(alpha=0.3), \n", + " nn.CELU(alpha=0.3),\n", " )\n", - " \n", + "\n", " self.layer3 = nn.Sequential(\n", " Conv2d(fmap_factor * 4, fmap_factor * 8, kernel_size=conv_ksize, padding=conv_pad, stride=1, bias=False),\n", " nn.MaxPool2d(kernel_size=2),\n", " GhostBatchNorm(fmap_factor * 8, num_splits=gbn_num_splits, weight=False),\n", " nn.CELU(alpha=0.3),\n", - " IdentityResidualBlock(fmap_factor * 8, \n", - " conv_ksize=conv_ksize, conv_pad=conv_pad, \n", - " gbn_num_splits=gbn_num_splits)\n", + " IdentityResidualBlock(\n", + " fmap_factor * 8, conv_ksize=conv_ksize, conv_pad=conv_pad, gbn_num_splits=gbn_num_splits\n", + " ),\n", " )\n", - " \n", + "\n", " self.pool = nn.MaxPool2d(kernel_size=4)\n", - " \n", - " self.classifier = nn.Sequential(\n", - " nn.Flatten(),\n", - " nn.Linear(fmap_factor * 8, num_classes)\n", - " )\n", + "\n", + " self.classifier = nn.Sequential(nn.Flatten(), nn.Linear(fmap_factor * 8, num_classes))\n", " self.scale = torch.tensor(0.0625, requires_grad=False)\n", "\n", " def forward(self, x):\n", @@ -240,8 +250,7 @@ " x = self.layer3(x)\n", " x = self.pool(x)\n", " y = self.classifier(x)\n", - " return y * self.scale\n", - " " + " return y * self.scale" ] }, { @@ -265,11 +274,12 @@ " num_params = 1\n", " for s in p.shape:\n", " num_params *= s\n", - " if display_all_modules: print(f\"{n}: {num_params}\")\n", + " if display_all_modules:\n", + " print(f\"{n}: {num_params}\")\n", " total_num_params += num_params\n", " print(\"-\" * 50)\n", " print(f\"Total number of parameters: {total_num_params:.2e}\")\n", - " \n", + "\n", "\n", "print_num_params(model)" ] @@ -319,20 +329,24 @@ "from torchvision.datasets.cifar import CIFAR10\n", "\n", "\n", - "train_transform = Compose([\n", - " Pad(4),\n", - " RandomCrop(32),\n", - " RandomHorizontalFlip(),\n", - " ToTensor(), \n", - " Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n", - " RandomErasing(scale=(0.0625, 0.0625), ratio=(1.0, 1.0))\n", - "])\n", + "train_transform = Compose(\n", + " [\n", + " Pad(4),\n", + " RandomCrop(32),\n", + " RandomHorizontalFlip(),\n", + " ToTensor(),\n", + " Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n", + " RandomErasing(scale=(0.0625, 0.0625), ratio=(1.0, 1.0)),\n", + " ]\n", + ")\n", "\n", "\n", - "test_transform = Compose([\n", - " ToTensor(), \n", - " Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n", - "])\n", + "test_transform = Compose(\n", + " [\n", + " ToTensor(),\n", + " Normalize((0.485, 0.456, 0.406), (0.229, 0.224, 0.225)),\n", + " ]\n", + ")\n", "\n", "\n", "train_ds = CIFAR10(\"/tmp/cifar10\", train=True, download=True, transform=train_transform)\n", @@ -371,18 +385,17 @@ "\n", "\n", "class CriterionWithLabelSmoothing(nn.Module):\n", - " \n", " def __init__(self, criterion, alpha=0.2):\n", " super(CriterionWithLabelSmoothing, self).__init__()\n", " self.criterion = criterion\n", - " if self.criterion.reduction != 'none':\n", + " if self.criterion.reduction != \"none\":\n", " raise ValueError(\"Input criterion should have reduction equal none\")\n", " self.alpha = alpha\n", - " \n", + "\n", " def forward(self, logits, targets):\n", " loss = self.criterion(logits, targets)\n", " log_probs = torch.log_softmax(logits, dim=1)\n", - " klloss = -log_probs.mean(dim=1) \n", + " klloss = -log_probs.mean(dim=1)\n", " out = (1.0 - self.alpha) * loss + self.alpha * klloss\n", " return out.mean(dim=0)" ] @@ -394,17 +407,20 @@ "outputs": [], "source": [ "def get_criterion(alpha):\n", - " return CriterionWithLabelSmoothing(nn.CrossEntropyLoss(reduction='none'), alpha=0.2)\n", + " return CriterionWithLabelSmoothing(nn.CrossEntropyLoss(reduction=\"none\"), alpha=0.2)\n", "\n", "\n", "def get_optimizer(model, momentum, weight_decay, nesterov):\n", " biases = [p for n, p in model.named_parameters() if \"bias\" in n]\n", " others = [p for n, p in model.named_parameters() if \"bias\" not in n]\n", " return optim.SGD(\n", - " [{\"params\": others, \"lr\": 1.0, \"weight_decay\": weight_decay}, \n", - " {\"params\": biases, \"lr\": 1.0, \"weight_decay\": weight_decay / 64}], \n", - " momentum=momentum, nesterov=nesterov\n", - " )\n" + " [\n", + " {\"params\": others, \"lr\": 1.0, \"weight_decay\": weight_decay},\n", + " {\"params\": biases, \"lr\": 1.0, \"weight_decay\": weight_decay / 64},\n", + " ],\n", + " momentum=momentum,\n", + " nesterov=nesterov,\n", + " )" ] }, { @@ -440,24 +456,23 @@ "source": [ "def get_lr_scheduler(optimizer, lr_max_value, lr_max_value_epoch, num_epochs, epoch_length):\n", " milestones_values = [\n", - " (0, 0.0), \n", - " (epoch_length * lr_max_value_epoch, lr_max_value), \n", - " (epoch_length * num_epochs - 1, 0.0)\n", + " (0, 0.0),\n", + " (epoch_length * lr_max_value_epoch, lr_max_value),\n", + " (epoch_length * num_epochs - 1, 0.0),\n", " ]\n", " lr_scheduler1 = PiecewiseLinear(optimizer, \"lr\", milestones_values=milestones_values, param_group_index=0)\n", "\n", " milestones_values = [\n", - " (0, 0.0), \n", - " (epoch_length * lr_max_value_epoch, lr_max_value * 64), \n", - " (epoch_length * num_epochs - 1, 0.0)\n", + " (0, 0.0),\n", + " (epoch_length * lr_max_value_epoch, lr_max_value * 64),\n", + " (epoch_length * num_epochs - 1, 0.0),\n", " ]\n", " lr_scheduler2 = PiecewiseLinear(optimizer, \"lr\", milestones_values=milestones_values, param_group_index=1)\n", "\n", " lr_scheduler = ParamGroupScheduler(\n", - " [lr_scheduler1, lr_scheduler2],\n", - " [\"lr scheduler (non-biases)\", \"lr scheduler (biases)\"]\n", + " [lr_scheduler1, lr_scheduler2], [\"lr scheduler (non-biases)\", \"lr scheduler (biases)\"]\n", " )\n", - " \n", + "\n", " return lr_scheduler" ] }, @@ -538,8 +553,12 @@ "from ignite.engine import create_supervised_trainer, create_supervised_evaluator, Events, convert_tensor\n", "from ignite.metrics import Accuracy\n", "from ignite.handlers import TensorboardLogger, ProgressBar\n", - "from ignite.handlers.tensorboard_logger import OutputHandler, OptimizerParamsHandler, GradsHistHandler, \\\n", - " global_step_from_engine" + "from ignite.handlers.tensorboard_logger import (\n", + " OutputHandler,\n", + " OptimizerParamsHandler,\n", + " GradsHistHandler,\n", + " global_step_from_engine,\n", + ")" ] }, { @@ -551,8 +570,10 @@ "# Transfer batch to GPU and set floating-point 16\n", "def prepare_batch_fp16(batch, device=None, non_blocking=True):\n", " x, y = batch\n", - " return (convert_tensor(x, device=device, non_blocking=non_blocking).half(),\n", - " convert_tensor(y, device=device, non_blocking=non_blocking))" + " return (\n", + " convert_tensor(x, device=device, non_blocking=non_blocking).half(),\n", + " convert_tensor(y, device=device, non_blocking=non_blocking),\n", + " )" ] }, { @@ -574,76 +595,80 @@ "\n", "\n", "def run_experiment(parameters):\n", - " device = 'cuda'\n", + " device = \"cuda\"\n", " fast_mode = parameters.get(\"fast_mode\", True)\n", - " \n", + "\n", " # setup model\n", - " model = FastResNet(\n", - " num_classes=10, \n", - " fmap_factor=parameters.get(\"fmap_factor\"), \n", - " conv_ksize=parameters.get(\"conv_ksize\"),\n", - " classif_scale=parameters.get(\"classif_scale\")\n", - " ).to(device).half()\n", - " \n", - " # setup dataloaders \n", + " model = (\n", + " FastResNet(\n", + " num_classes=10,\n", + " fmap_factor=parameters.get(\"fmap_factor\"),\n", + " conv_ksize=parameters.get(\"conv_ksize\"),\n", + " classif_scale=parameters.get(\"classif_scale\"),\n", + " )\n", + " .to(device)\n", + " .half()\n", + " )\n", + "\n", + " # setup dataloaders\n", " train_loader, test_loader = get_train_test_loaders()\n", - " \n", + "\n", " # setup solver\n", " criterion = get_criterion(parameters.get(\"alpha\")).to(device)\n", " optimizer = get_optimizer(\n", - " model, \n", - " parameters.get(\"momentum\"), \n", - " parameters.get(\"weight_decay\"),\n", - " parameters.get(\"nesterov\")\n", + " model, parameters.get(\"momentum\"), parameters.get(\"weight_decay\"), parameters.get(\"nesterov\")\n", " )\n", " lr_scheduler = get_lr_scheduler(\n", - " optimizer, \n", + " optimizer,\n", " parameters.get(\"lr_max_value\"),\n", - " parameters.get(\"lr_max_value_epoch\"), \n", + " parameters.get(\"lr_max_value_epoch\"),\n", " num_epochs=num_epochs,\n", - " epoch_length=len(train_loader)\n", + " epoch_length=len(train_loader),\n", " )\n", - " \n", + "\n", " # setup ignite trainer\n", - " trainer = create_supervised_trainer(model, optimizer, criterion, \n", - " device=device, non_blocking=True,\n", - " prepare_batch=prepare_batch_fp16)\n", - " \n", + " trainer = create_supervised_trainer(\n", + " model, optimizer, criterion, device=device, non_blocking=True, prepare_batch=prepare_batch_fp16\n", + " )\n", + "\n", " # setup learning rate scheduler\n", " trainer.add_event_handler(Events.ITERATION_STARTED, lr_scheduler)\n", - " \n", + "\n", " # setup tensorboard logger\n", - " exp_log_name = f\"exp_{parameters.get('fmap_factor')}_{parameters.get('conv_ksize')}_\" + \\\n", - " f\"{parameters.get('alpha'):.2}_{parameters.get('lr_max_value'):.4}\"\n", + " exp_log_name = (\n", + " f\"exp_{parameters.get('fmap_factor')}_{parameters.get('conv_ksize')}_\"\n", + " + f\"{parameters.get('alpha'):.2}_{parameters.get('lr_max_value'):.4}\"\n", + " )\n", " tb_logger = TensorboardLogger(log_dir=f\"/tmp/tb_logs/{exp_log_name}\")\n", - " \n", + "\n", " if not fast_mode:\n", " # - log learning rate\n", " tb_logger.attach(trainer, OptimizerParamsHandler(optimizer), event_name=Events.ITERATION_STARTED)\n", "\n", " # - log training batch loss\n", - " tb_logger.attach(trainer, OutputHandler(tag=\"training\", output_transform=lambda x: {\"batch loss\": x}), \n", - " event_name=Events.ITERATION_COMPLETED)\n", + " tb_logger.attach(\n", + " trainer,\n", + " OutputHandler(tag=\"training\", output_transform=lambda x: {\"batch loss\": x}),\n", + " event_name=Events.ITERATION_COMPLETED,\n", + " )\n", "\n", " # - log model grads\n", - " tb_logger.attach(trainer, GradsHistHandler(model), event_name=Events.EPOCH_COMPLETED) \n", - " \n", + " tb_logger.attach(trainer, GradsHistHandler(model), event_name=Events.EPOCH_COMPLETED)\n", + "\n", " # setup a progress bar\n", - " ProgressBar().attach(trainer, event_name=Events.EPOCH_COMPLETED, closing_event_name=Events.COMPLETED) \n", - " \n", + " ProgressBar().attach(trainer, event_name=Events.EPOCH_COMPLETED, closing_event_name=Events.COMPLETED)\n", + "\n", " # setup evaluator\n", " def output_transform(output):\n", " y_pred, y = output\n", " y_pred = y_pred.float()\n", " return y_pred, y\n", "\n", - " metrics = {\n", - " \"test accuracy\": Accuracy(output_transform=output_transform)\n", - " }\n", - " evaluator = create_supervised_evaluator(model, metrics=metrics, \n", - " device=device, non_blocking=True, \n", - " prepare_batch=prepare_batch_fp16)\n", - " \n", + " metrics = {\"test accuracy\": Accuracy(output_transform=output_transform)}\n", + " evaluator = create_supervised_evaluator(\n", + " model, metrics=metrics, device=device, non_blocking=True, prepare_batch=prepare_batch_fp16\n", + " )\n", + "\n", " # evaluate trained model each 3 epochs\n", " @trainer.on(Events.EPOCH_COMPLETED)\n", " def run_evaluation(engine):\n", @@ -651,21 +676,22 @@ " c2 = engine.state.epoch == engine.state.max_epochs\n", " if (c1 and not fast_mode) or c2:\n", " evaluator.run(test_loader)\n", - " \n", + "\n", " if not fast_mode:\n", " # - log test accuracy\n", - " tb_logger.attach(evaluator, \n", - " OutputHandler(tag=\"validation\", metric_names=\"all\", \n", - " global_step_transform=global_step_from_engine(trainer)), \n", - " event_name=Events.EPOCH_COMPLETED)\n", - "\n", - " trainer.run(train_loader, max_epochs=num_epochs) \n", - " test_acc = evaluator.state.metrics['test accuracy']\n", - " \n", + " tb_logger.attach(\n", + " evaluator,\n", + " OutputHandler(tag=\"validation\", metric_names=\"all\", global_step_transform=global_step_from_engine(trainer)),\n", + " event_name=Events.EPOCH_COMPLETED,\n", + " )\n", + "\n", + " trainer.run(train_loader, max_epochs=num_epochs)\n", + " test_acc = evaluator.state.metrics[\"test accuracy\"]\n", + "\n", " # dump hparams/result to Tensorboard\n", - " tb_logger.writer.add_hparams(parameters, {'hparam/test_accuracy': test_acc})\n", + " tb_logger.writer.add_hparams(parameters, {\"hparam/test_accuracy\": test_acc})\n", "\n", - " tb_logger.close() \n", + " tb_logger.close()\n", " return test_acc" ] }, @@ -696,7 +722,7 @@ " \"nesterov\": True,\n", " \"lr_max_value\": 1.0,\n", " \"lr_max_value_epoch\": num_epochs // 5,\n", - " \"fast_mode\": False\n", + " \"fast_mode\": False,\n", " }\n", ")" ] @@ -745,7 +771,7 @@ " \"type\": \"range\",\n", " \"bounds\": [1e-4, 1e-3],\n", " \"value_type\": \"float\",\n", - " }, \n", + " },\n", " {\n", " \"name\": \"nesterov\",\n", " \"type\": \"choice\",\n", @@ -761,7 +787,7 @@ " \"type\": \"range\",\n", " \"bounds\": [1, 10],\n", " },\n", - "]\n" + "]" ] }, { @@ -781,11 +807,8 @@ "\n", "\n", "best_parameters, values, experiment, model = optimize(\n", - " parameters=parameters_space,\n", - " evaluation_function=run_experiment,\n", - " objective_name='test accuracy',\n", - " total_trials=30\n", - ")\n" + " parameters=parameters_space, evaluation_function=run_experiment, objective_name=\"test accuracy\", total_trials=30\n", + ")" ] }, { @@ -819,7 +842,7 @@ "metadata": {}, "outputs": [], "source": [ - "render(plot_contour(model=model, param_x='lr_max_value', param_y='momentum', metric_name='test accuracy'))" + "render(plot_contour(model=model, param_x=\"lr_max_value\", param_y=\"momentum\", metric_name=\"test accuracy\"))" ] }, { @@ -839,11 +862,9 @@ "num_epochs = 20\n", "\n", "best_parameters_copy = dict(best_parameters)\n", - "best_parameters_copy['fast_mode'] = False\n", + "best_parameters_copy[\"fast_mode\"] = False\n", "\n", - "run_experiment(\n", - " parameters=best_parameters_copy\n", - ")" + "run_experiment(parameters=best_parameters_copy)" ] }, { diff --git a/examples/notebooks/CycleGAN_with_nvidia_apex.ipynb b/examples/notebooks/CycleGAN_with_nvidia_apex.ipynb index 196e4417714d..bc8dc31f2c28 100644 --- a/examples/notebooks/CycleGAN_with_nvidia_apex.ipynb +++ b/examples/notebooks/CycleGAN_with_nvidia_apex.ipynb @@ -86,6 +86,7 @@ "outputs": [], "source": [ "import torch\n", + "\n", "torch.__version__" ] }, @@ -132,6 +133,7 @@ "outputs": [], "source": [ "import ignite\n", + "\n", "ignite.__version__" ] }, @@ -167,19 +169,19 @@ "from torch.utils.data import Dataset, DataLoader\n", "from PIL import Image\n", "\n", + "\n", "class FilesDataset(Dataset):\n", - " \n", " def __init__(self, path, extension=\"*.jpg\"):\n", " self.path = Path(path)\n", " assert self.path.exists(), \"Path '{}' is not found\".format(path)\n", " self.images = list(self.path.rglob(extension))\n", " assert len(self.images) > 0, \"No images with extension {} found at '{}'\".format(extension, path)\n", - " \n", + "\n", " def __len__(self):\n", " return len(self.images)\n", - " \n", + "\n", " def __getitem__(self, i):\n", - " return Image.open(self.images[i]).convert('RGB')" + " return Image.open(self.images[i]).convert(\"RGB\")" ] }, { @@ -195,7 +197,7 @@ "train_A = FilesDataset(root / \"trainA\")\n", "train_B = FilesDataset(root / \"trainB\")\n", "\n", - "test_A = FilesDataset(root / \"testA\") \n", + "test_A = FilesDataset(root / \"testA\")\n", "test_B = FilesDataset(root / \"testB\")" ] }, @@ -212,7 +214,11 @@ "metadata": {}, "outputs": [], "source": [ - "print(\"Dataset sizes: \\ntrain A: {} | B: {}\\ntest A: {} | B: {}\\n\\t\".format(len(train_A), len(train_B), len(test_A), len(test_B)))" + "print(\n", + " \"Dataset sizes: \\ntrain A: {} | B: {}\\ntest A: {} | B: {}\\n\\t\".format(\n", + " len(train_A), len(train_B), len(test_A), len(test_B)\n", + " )\n", + ")" ] }, { @@ -240,6 +246,7 @@ "outputs": [], "source": [ "import matplotlib.pylab as plt\n", + "\n", "%matplotlib inline" ] }, @@ -290,11 +297,10 @@ "\n", "\n", "class Image2ImageDataset(Dataset):\n", - " \n", " def __init__(self, ds_a, ds_b):\n", " self.dataset_a = ds_a\n", " self.dataset_b = ds_b\n", - " \n", + "\n", " def __len__(self):\n", " return max(len(self.dataset_a), len(self.dataset_b))\n", "\n", @@ -302,21 +308,17 @@ " dp_a = self.dataset_a[i % len(self.dataset_a)]\n", " j = random.randint(0, len(self.dataset_b) - 1)\n", " dp_b = self.dataset_b[j]\n", - " return {\n", - " 'A': dp_a,\n", - " 'B': dp_b\n", - " }\n", + " return {\"A\": dp_a, \"B\": dp_b}\n", "\n", "\n", "class TransformedDataset(Dataset):\n", - " \n", " def __init__(self, ds, transform):\n", " self.dataset = ds\n", " self.transform = transform\n", - " \n", + "\n", " def __len__(self):\n", " return len(self.dataset)\n", - " \n", + "\n", " def __getitem__(self, i):\n", " return {k: self.transform(v) for k, v in self.dataset[i].items()}" ] @@ -342,10 +344,10 @@ "plt.figure(figsize=(10, 5))\n", "plt.subplot(121)\n", "plt.title(\"Train dataset 'Horses'\")\n", - "plt.imshow(dp['A'])\n", + "plt.imshow(dp[\"A\"])\n", "plt.subplot(122)\n", "plt.title(\"Train dataset 'Zebras'\")\n", - "plt.imshow(dp['B'])" + "plt.imshow(dp[\"B\"])" ] }, { @@ -359,10 +361,10 @@ "plt.figure(figsize=(10, 5))\n", "plt.subplot(121)\n", "plt.title(\"Test dataset 'Horses'\")\n", - "plt.imshow(dp['A'])\n", + "plt.imshow(dp[\"A\"])\n", "plt.subplot(122)\n", "plt.title(\"Test dataset 'Zebras'\")\n", - "plt.imshow(dp['B'])" + "plt.imshow(dp[\"B\"])" ] }, { @@ -374,27 +376,30 @@ "from torchvision.transforms import Compose, ColorJitter, RandomHorizontalFlip, ToTensor, Normalize, RandomCrop\n", "\n", "# To accelerate the training we reduce the image size to 200x200 pix instead of 256x256\n", - "train_transform = Compose([\n", - " RandomCrop(200),\n", - " RandomHorizontalFlip(),\n", - " ColorJitter(),\n", - " ToTensor(),\n", - " Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))\n", - "])\n", + "train_transform = Compose(\n", + " [\n", + " RandomCrop(200),\n", + " RandomHorizontalFlip(),\n", + " ColorJitter(),\n", + " ToTensor(),\n", + " Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)),\n", + " ]\n", + ")\n", "transformed_train_ab_ds = TransformedDataset(train_ab_ds, transform=train_transform)\n", "\n", "# Please select appropriate batch_size value according to your infrastructure\n", "batch_size = 10\n", - "train_ab_loader = DataLoader(transformed_train_ab_ds, batch_size=batch_size, shuffle=True, drop_last=True, pin_memory=True, num_workers=4)\n", + "train_ab_loader = DataLoader(\n", + " transformed_train_ab_ds, batch_size=batch_size, shuffle=True, drop_last=True, pin_memory=True, num_workers=4\n", + ")\n", "\n", "\n", - "test_transform = Compose([\n", - " ToTensor(),\n", - " Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))\n", - "])\n", + "test_transform = Compose([ToTensor(), Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))])\n", "transformed_test_ab_ds = TransformedDataset(test_ab_ds, transform=test_transform)\n", "batch_size = 10\n", - "test_ab_loader = DataLoader(transformed_test_ab_ds, batch_size=batch_size, shuffle=False, drop_last=False, pin_memory=True, num_workers=4)" + "test_ab_loader = DataLoader(\n", + " transformed_test_ab_ds, batch_size=batch_size, shuffle=False, drop_last=False, pin_memory=True, num_workers=4\n", + ")" ] }, { @@ -411,16 +416,12 @@ "plt.figure(figsize=(16, 8))\n", "plt.axis(\"off\")\n", "plt.title(\"Training Images from A\")\n", - "plt.imshow( \n", - " vutils.make_grid(real_batch['A'][:64], padding=2, normalize=True).cpu().numpy().transpose((1, 2, 0))\n", - ")\n", + "plt.imshow(vutils.make_grid(real_batch[\"A\"][:64], padding=2, normalize=True).cpu().numpy().transpose((1, 2, 0)))\n", "\n", "plt.figure(figsize=(16, 8))\n", "plt.axis(\"off\")\n", "plt.title(\"Training Images from B\")\n", - "plt.imshow(\n", - " vutils.make_grid(real_batch['B'][:64], padding=2, normalize=True).cpu().numpy().transpose((1, 2, 0))\n", - ")\n", + "plt.imshow(vutils.make_grid(real_batch[\"B\"][:64], padding=2, normalize=True).cpu().numpy().transpose((1, 2, 0)))\n", "real_batch = None\n", "torch.cuda.empty_cache()" ] @@ -468,29 +469,27 @@ " return nn.Sequential(\n", " nn.ConvTranspose2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride, padding=1, output_padding=1),\n", " nn.InstanceNorm2d(out_planes, affine=False, track_running_stats=False),\n", - " nn.ReLU(inplace=True)\n", + " nn.ReLU(inplace=True),\n", " )\n", "\n", "\n", "class ResidualBlock(nn.Module):\n", - " \n", " def __init__(self, in_planes):\n", " super(ResidualBlock, self).__init__()\n", " self.conv1 = get_conv_inorm_relu(in_planes, in_planes, kernel_size=3, stride=1)\n", - " self.conv2 = get_conv_inorm_relu(in_planes, in_planes, kernel_size=3, stride=1, with_relu=False) \n", + " self.conv2 = get_conv_inorm_relu(in_planes, in_planes, kernel_size=3, stride=1, with_relu=False)\n", "\n", " def forward(self, x):\n", " residual = x\n", " x = self.conv1(x)\n", - " x = self.conv2(x) \n", + " x = self.conv2(x)\n", " return x + residual\n", "\n", "\n", "class Generator(nn.Module):\n", - " \n", " def __init__(self):\n", " super(Generator, self).__init__()\n", - " \n", + "\n", " self.c7s1_64 = get_conv_inorm_relu(3, 64, kernel_size=7, stride=1)\n", " self.d128 = get_conv_inorm_relu(64, 128, kernel_size=3, stride=2, reflection_pad=False)\n", " self.d256 = get_conv_inorm_relu(128, 256, kernel_size=3, stride=2, reflection_pad=False)\n", @@ -508,7 +507,7 @@ " x = self.c7s1_64(x)\n", " x = self.d128(x)\n", " x = self.d256(x)\n", - " \n", + "\n", " # 9 residual blocks\n", " x = self.resnet9(x)\n", "\n", @@ -516,7 +515,7 @@ " x = self.u128(x)\n", " x = self.u64(x)\n", " y = self.c7s1_3(x)\n", - " return y\n" + " return y" ] }, { @@ -571,12 +570,11 @@ " return nn.Sequential(\n", " nn.Conv2d(in_planes, out_planes, kernel_size=4, stride=stride, padding=1),\n", " nn.InstanceNorm2d(out_planes, affine=False, track_running_stats=False),\n", - " nn.LeakyReLU(negative_slope=negative_slope, inplace=True)\n", + " nn.LeakyReLU(negative_slope=negative_slope, inplace=True),\n", " )\n", "\n", "\n", "class Discriminator(nn.Module):\n", - "\n", " def __init__(self):\n", " super(Discriminator, self).__init__()\n", " self.c64 = nn.Conv2d(3, 64, kernel_size=4, stride=2, padding=1)\n", @@ -594,7 +592,7 @@ " x = self.c256(x)\n", " x = self.c512(x)\n", " y = self.last_conv(x)\n", - " return y\n" + " return y" ] }, { @@ -645,7 +643,8 @@ "metadata": {}, "outputs": [], "source": [ - "g = None; d = None" + "g = None\n", + "d = None" ] }, { @@ -735,9 +734,12 @@ "\n", "\n", "# Initialize Amp\n", - "models, optimizers = amp.initialize([generator_A2B, generator_B2A, discriminator_A, discriminator_B], \n", - " [optimizer_G, optimizer_D],\n", - " opt_level=\"O2\", num_losses=2)\n", + "models, optimizers = amp.initialize(\n", + " [generator_A2B, generator_B2A, discriminator_A, discriminator_B],\n", + " [optimizer_G, optimizer_D],\n", + " opt_level=\"O2\",\n", + " num_losses=2,\n", + ")\n", "\n", "generator_A2B, generator_B2A, discriminator_A, discriminator_B = models\n", "optimizer_G, optimizer_D = optimizers" @@ -773,7 +775,7 @@ " buffer.append(b.cpu())\n", " elif random.uniform(0, 1) > 0.5:\n", " # Add newly created image into the buffer and put ont from the buffer into the output\n", - " random_index = random.randint(0, buffer_size - 1) \n", + " random_index = random.randint(0, buffer_size - 1)\n", " output_batch.append(buffer[random_index].clone().to(device))\n", " buffer[random_index] = b.cpu()\n", " else:\n", @@ -811,7 +813,7 @@ "\n", "def discriminator_forward_pass(discriminator, batch_real, batch_fake, fake_buffer):\n", " decision_real = discriminator(batch_real)\n", - " batch_fake = buffer_insert_and_get(fake_buffer, batch_fake) \n", + " batch_fake = buffer_insert_and_get(fake_buffer, batch_fake)\n", " decision_fake = discriminator(batch_fake)\n", " return decision_real, decision_fake\n", "\n", @@ -821,12 +823,12 @@ " target = torch.ones_like(batch_decision)\n", " loss_gan = F.mse_loss(batch_decision, target)\n", " # loss cycle\n", - " loss_cycle = F.l1_loss(batch_rec, batch_real) * lambda_value \n", + " loss_cycle = F.l1_loss(batch_rec, batch_real) * lambda_value\n", " return loss_gan + loss_cycle\n", "\n", "\n", "def compute_loss_discriminator(decision_real, decision_fake):\n", - " # loss = mean (D_b(y) − 1)^2 + mean D_b(G(x))^2 \n", + " # loss = mean (D_b(y) − 1)^2 + mean D_b(G(x))^2\n", " loss = F.mse_loss(decision_fake, torch.zeros_like(decision_fake))\n", " loss += F.mse_loss(decision_real, torch.ones_like(decision_real))\n", " return loss\n", @@ -838,15 +840,15 @@ " discriminator_A.train()\n", " discriminator_B.train()\n", "\n", - " real_a = convert_tensor(batch['A'], device=device, non_blocking=True)\n", - " real_b = convert_tensor(batch['B'], device=device, non_blocking=True)\n", - " \n", + " real_a = convert_tensor(batch[\"A\"], device=device, non_blocking=True)\n", + " real_b = convert_tensor(batch[\"B\"], device=device, non_blocking=True)\n", + "\n", " # Update generators:\n", "\n", " # Disable grads computation for the discriminators:\n", " toggle_grad(discriminator_A, False)\n", - " toggle_grad(discriminator_B, False) \n", - " \n", + " toggle_grad(discriminator_B, False)\n", + "\n", " fake_b = generator_A2B(real_a)\n", " rec_a = generator_B2A(fake_b)\n", " fake_a = generator_B2A(real_b)\n", @@ -855,8 +857,8 @@ " decision_fake_b = discriminator_B(fake_b)\n", "\n", " # Compute loss for generators and update generators\n", - " # loss_a2b = GAN loss: mean (D_B(G(x)) − 1)^2 + Forward cycle loss: || F(G(x)) - x ||_1 \n", - " loss_a2b = compute_loss_generator(decision_fake_b, real_a, rec_a, lambda_value) \n", + " # loss_a2b = GAN loss: mean (D_B(G(x)) − 1)^2 + Forward cycle loss: || F(G(x)) - x ||_1\n", + " loss_a2b = compute_loss_generator(decision_fake_b, real_a, rec_a, lambda_value)\n", "\n", " # loss_b2a = GAN loss: mean (D_A(F(x)) − 1)^2 + Backward cycle loss: || G(F(y)) - y ||_1\n", " loss_b2a = compute_loss_generator(decision_fake_a, real_b, rec_b, lambda_value)\n", @@ -864,36 +866,40 @@ " # total generators loss:\n", " loss_generators = loss_a2b + loss_b2a\n", "\n", - " optimizer_G.zero_grad() \n", + " optimizer_G.zero_grad()\n", " with amp.scale_loss(loss_generators, optimizer_G, loss_id=0) as scaled_loss:\n", " scaled_loss.backward()\n", " optimizer_G.step()\n", "\n", " decision_fake_a = rec_a = decision_fake_b = rec_b = None\n", - " \n", + "\n", " # Update discriminators:\n", "\n", " # Enable grads computation for the discriminators:\n", " toggle_grad(discriminator_A, True)\n", " toggle_grad(discriminator_B, True)\n", "\n", - " decision_real_a, decision_fake_a = discriminator_forward_pass(discriminator_A, real_a, fake_a.detach(), fake_a_buffer) \n", - " decision_real_b, decision_fake_b = discriminator_forward_pass(discriminator_B, real_b, fake_b.detach(), fake_b_buffer) \n", + " decision_real_a, decision_fake_a = discriminator_forward_pass(\n", + " discriminator_A, real_a, fake_a.detach(), fake_a_buffer\n", + " )\n", + " decision_real_b, decision_fake_b = discriminator_forward_pass(\n", + " discriminator_B, real_b, fake_b.detach(), fake_b_buffer\n", + " )\n", " # Compute loss for discriminators and update discriminators\n", " # loss_a = mean (D_a(y) − 1)^2 + mean D_a(F(x))^2\n", " loss_a = compute_loss_discriminator(decision_real_a, decision_fake_a)\n", "\n", " # loss_b = mean (D_b(y) − 1)^2 + mean D_b(G(x))^2\n", " loss_b = compute_loss_discriminator(decision_real_b, decision_fake_b)\n", - " \n", + "\n", " # total discriminators loss:\n", " loss_discriminators = 0.5 * (loss_a + loss_b)\n", - " \n", + "\n", " optimizer_D.zero_grad()\n", " with amp.scale_loss(loss_discriminators, optimizer_D, loss_id=1) as scaled_loss:\n", " scaled_loss.backward()\n", " optimizer_D.step()\n", - " \n", + "\n", " return {\n", " \"loss_generators\": loss_generators.item(),\n", " \"loss_generator_a2b\": loss_a2b.item(),\n", @@ -901,8 +907,7 @@ " \"loss_discriminators\": loss_discriminators.item(),\n", " \"loss_discriminator_a\": loss_a.item(),\n", " \"loss_discriminator_b\": loss_b.item(),\n", - " }\n", - " " + " }" ] }, { @@ -981,20 +986,22 @@ "trainer = Engine(update_fn)\n", "\n", "metric_names = [\n", - " 'loss_discriminators', \n", - " 'loss_generators', \n", - " 'loss_discriminator_a',\n", - " 'loss_discriminator_b',\n", - " 'loss_generator_a2b',\n", - " 'loss_generator_b2a' \n", + " \"loss_discriminators\",\n", + " \"loss_generators\",\n", + " \"loss_discriminator_a\",\n", + " \"loss_discriminator_b\",\n", + " \"loss_generator_a2b\",\n", + " \"loss_generator_b2a\",\n", "]\n", "\n", + "\n", "def output_transform(out, name):\n", " return out[name]\n", "\n", + "\n", "for name in metric_names:\n", " # here we cannot use lambdas as they do not store argument `name`\n", - " RunningAverage(output_transform=partial(output_transform, name=name)).attach(trainer, name)\n" + " RunningAverage(output_transform=partial(output_transform, name=name)).attach(trainer, name)" ] }, { @@ -1008,9 +1015,7 @@ "exp_name = datetime.now().strftime(\"%Y%m%d-%H%M%S\")\n", "tb_logger = TensorboardLogger(log_dir=\"/tmp/cycle_gan_horse2zebra_tb_logs/{}\".format(exp_name))\n", "\n", - "tb_logger.attach(trainer, \n", - " log_handler=OutputHandler('training', metric_names), \n", - " event_name=Events.ITERATION_COMPLETED)\n", + "tb_logger.attach(trainer, log_handler=OutputHandler(\"training\", metric_names), event_name=Events.ITERATION_COMPLETED)\n", "\n", "print(\"Experiment name: \", exp_name)" ] @@ -1024,17 +1029,12 @@ "from pathlib import Path\n", "\n", "try:\n", - "\n", " wb_run_name = \"cycle_gan_horse2zebra\"\n", " wb_dir = Path(\"/tmp/cycle_gan_horse2zebra_wandb\")\n", " if not wb_dir.exists():\n", " wb_dir.mkdir()\n", " wb_logger = WandBLogger(\n", - " project=\"ignite-cyclegan-apex\",\n", - " name=wb_run_name,\n", - " sync_tensorboard=True,\n", - " dir=wb_dir.as_posix(),\n", - " reinit=True\n", + " project=\"ignite-cyclegan-apex\", name=wb_run_name, sync_tensorboard=True, dir=wb_dir.as_posix(), reinit=True\n", " )\n", "except RuntimeError:\n", " wb_logger = None" @@ -1058,24 +1058,24 @@ "\n", "def evaluate_fn(engine, batch):\n", " generator_A2B.eval()\n", - " generator_B2A.eval() \n", + " generator_B2A.eval()\n", " with torch.no_grad():\n", - " real_a = convert_tensor(batch['A'], device=device, non_blocking=True)\n", - " real_b = convert_tensor(batch['B'], device=device, non_blocking=True)\n", - " \n", + " real_a = convert_tensor(batch[\"A\"], device=device, non_blocking=True)\n", + " real_b = convert_tensor(batch[\"B\"], device=device, non_blocking=True)\n", + "\n", " fake_b = generator_A2B(real_a)\n", " rec_a = generator_B2A(fake_b)\n", "\n", " fake_a = generator_B2A(real_b)\n", " rec_b = generator_A2B(fake_a)\n", - " \n", + "\n", " return {\n", - " 'real_a': real_a,\n", - " 'real_b': real_b,\n", - " 'fake_a': fake_a,\n", - " 'fake_b': fake_b,\n", - " 'rec_a': rec_a,\n", - " 'rec_b': rec_b, \n", + " \"real_a\": real_a,\n", + " \"real_b\": real_b,\n", + " \"fake_a\": fake_a,\n", + " \"fake_b\": fake_b,\n", + " \"rec_a\": rec_a,\n", + " \"rec_b\": rec_b,\n", " }\n", "\n", "\n", @@ -1100,8 +1100,12 @@ "small_test_ds = Subset(test_ab_ds, test_random_indices)\n", "small_test_ds = TransformedDataset(small_test_ds, transform=test_transform)\n", "\n", - "eval_train_loader = DataLoader(small_train_ds, batch_size=eval_batch_size, shuffle=False, drop_last=False, pin_memory=True, num_workers=4)\n", - "eval_test_loader = DataLoader(small_test_ds, batch_size=eval_batch_size, shuffle=False, drop_last=False, pin_memory=True, num_workers=4)" + "eval_train_loader = DataLoader(\n", + " small_train_ds, batch_size=eval_batch_size, shuffle=False, drop_last=False, pin_memory=True, num_workers=4\n", + ")\n", + "eval_test_loader = DataLoader(\n", + " small_test_ds, batch_size=eval_batch_size, shuffle=False, drop_last=False, pin_memory=True, num_workers=4\n", + ")" ] }, { @@ -1117,7 +1121,6 @@ "\n", "\n", "def log_generated_images(engine, logger, event_name):\n", - "\n", " tag = \"Train\" if engine.state.dataloader == eval_train_loader else \"Test\"\n", " output = engine.state.output\n", " state = trainer.state\n", @@ -1127,30 +1130,50 @@ " # [real a1, real a2, ...]\n", " # [fake a1, fake a2, ...]\n", " # [rec a1, rec a2, ...]\n", - " \n", - " s = output['real_a'].shape[0]\n", - " res_a = vutils.make_grid(torch.cat([\n", - " output['real_a'],\n", - " output['fake_b'],\n", - " output['rec_a'],\n", - " ]), padding=2, normalize=True, nrow=s).cpu()\n", "\n", - " logger.writer.add_image(tag=\"{} Horses2Zebras (real, fake, rec)\".format(tag), \n", - " img_tensor=res_a, global_step=global_step, dataformats='CHW')\n", + " s = output[\"real_a\"].shape[0]\n", + " res_a = vutils.make_grid(\n", + " torch.cat(\n", + " [\n", + " output[\"real_a\"],\n", + " output[\"fake_b\"],\n", + " output[\"rec_a\"],\n", + " ]\n", + " ),\n", + " padding=2,\n", + " normalize=True,\n", + " nrow=s,\n", + " ).cpu()\n", + "\n", + " logger.writer.add_image(\n", + " tag=\"{} Horses2Zebras (real, fake, rec)\".format(tag),\n", + " img_tensor=res_a,\n", + " global_step=global_step,\n", + " dataformats=\"CHW\",\n", + " )\n", + "\n", + " s = output[\"real_b\"].shape[0]\n", + " res_b = vutils.make_grid(\n", + " torch.cat(\n", + " [\n", + " output[\"real_b\"],\n", + " output[\"fake_a\"],\n", + " output[\"rec_b\"],\n", + " ]\n", + " ),\n", + " padding=2,\n", + " normalize=True,\n", + " nrow=s,\n", + " ).cpu()\n", + " logger.writer.add_image(\n", + " tag=\"{} Zebras2Horses (real, fake, rec)\".format(tag),\n", + " img_tensor=res_b,\n", + " global_step=global_step,\n", + " dataformats=\"CHW\",\n", + " )\n", "\n", - " s = output['real_b'].shape[0]\n", - " res_b = vutils.make_grid(torch.cat([\n", - " output['real_b'],\n", - " output['fake_a'],\n", - " output['rec_b'],\n", - " ]), padding=2, normalize=True, nrow=s).cpu()\n", - " logger.writer.add_image(tag=\"{} Zebras2Horses (real, fake, rec)\".format(tag), \n", - " img_tensor=res_b, global_step=global_step, dataformats='CHW')\n", "\n", - " \n", - "tb_logger.attach(evaluator,\n", - " log_handler=log_generated_images, \n", - " event_name=Events.COMPLETED)" + "tb_logger.attach(evaluator, log_handler=log_generated_images, event_name=Events.COMPLETED)" ] }, { @@ -1170,23 +1193,18 @@ "\n", "lr = 0.0002\n", "\n", - "milestones_values = [\n", - " (0, lr),\n", - " (100, lr),\n", - " (200, 0.0)\n", - "]\n", - "gen_lr_scheduler = PiecewiseLinear(optimizer_D, param_name='lr', milestones_values=milestones_values)\n", - "desc_lr_scheduler = PiecewiseLinear(optimizer_G, param_name='lr', milestones_values=milestones_values)\n", + "milestones_values = [(0, lr), (100, lr), (200, 0.0)]\n", + "gen_lr_scheduler = PiecewiseLinear(optimizer_D, param_name=\"lr\", milestones_values=milestones_values)\n", + "desc_lr_scheduler = PiecewiseLinear(optimizer_G, param_name=\"lr\", milestones_values=milestones_values)\n", "\n", - "lr_scheduler = ParamGroupScheduler([gen_lr_scheduler, desc_lr_scheduler], \n", - " names=['gen_lr_scheduler', 'desc_lr_scheduler'])\n", + "lr_scheduler = ParamGroupScheduler(\n", + " [gen_lr_scheduler, desc_lr_scheduler], names=[\"gen_lr_scheduler\", \"desc_lr_scheduler\"]\n", + ")\n", "\n", "trainer.add_event_handler(Events.EPOCH_STARTED, lr_scheduler)\n", "\n", "\n", - "tb_logger.attach(trainer,\n", - " log_handler=OptimizerParamsHandler(optimizer_G, \"lr\"), \n", - " event_name=Events.EPOCH_STARTED)" + "tb_logger.attach(trainer, log_handler=OptimizerParamsHandler(optimizer_G, \"lr\"), event_name=Events.EPOCH_STARTED)" ] }, { @@ -1224,7 +1242,6 @@ " \"discriminator_B\": discriminator_B,\n", " \"generator_B2A\": generator_B2A,\n", " \"discriminator_A\": discriminator_A,\n", - " \n", " \"optimizer_G\": optimizer_G,\n", " \"optimizer_D\": optimizer_D,\n", "}\n", @@ -1244,8 +1261,12 @@ "# Iteration-wise progress bar\n", "ProgressBar(bar_format=\"\").attach(trainer)\n", "# Epoch-wise progress bar with display of training losses\n", - "ProgressBar(persist=True, bar_format=\"\").attach(trainer, metric_names=['loss_discriminators', 'loss_generators'], \n", - " event_name=Events.EPOCH_STARTED, closing_event_name=Events.COMPLETED)" + "ProgressBar(persist=True, bar_format=\"\").attach(\n", + " trainer,\n", + " metric_names=[\"loss_discriminators\", \"loss_generators\"],\n", + " event_name=Events.EPOCH_STARTED,\n", + " closing_event_name=Events.COMPLETED,\n", + ")" ] }, { @@ -1340,16 +1361,16 @@ "outputs": [], "source": [ "i = random.randint(0, len(test_ab_ds) - 1)\n", - "img = test_ab_ds[i]['A']\n", + "img = test_ab_ds[i][\"A\"]\n", "x = test_transform(img)\n", "x = x.unsqueeze(0).to(device)\n", "\n", "\n", "with torch.no_grad():\n", " y_pred = generator_A2B(x)\n", - " \n", "\n", - "img_pred = (255 * normalize(y_pred[0, ...])).cpu().numpy().transpose((1, 2, 0)).astype('uint8')" + "\n", + "img_pred = (255 * normalize(y_pred[0, ...])).cpu().numpy().transpose((1, 2, 0)).astype(\"uint8\")" ] }, { @@ -1400,7 +1421,7 @@ " y_pred = generator_A2B(x)\n", "\n", "\n", - "img_pred = (255 * normalize(y_pred[0, ...])).cpu().numpy().transpose((1, 2, 0)).astype('uint8')" + "img_pred = (255 * normalize(y_pred[0, ...])).cpu().numpy().transpose((1, 2, 0)).astype(\"uint8\")" ] }, { diff --git a/examples/notebooks/CycleGAN_with_torch_cuda_amp.ipynb b/examples/notebooks/CycleGAN_with_torch_cuda_amp.ipynb index df53350ac2c4..9ddfa2d4b1e4 100644 --- a/examples/notebooks/CycleGAN_with_torch_cuda_amp.ipynb +++ b/examples/notebooks/CycleGAN_with_torch_cuda_amp.ipynb @@ -109,6 +109,7 @@ "source": [ "import torch\n", "import ignite\n", + "\n", "torch.__version__, ignite.__version__" ] }, @@ -144,19 +145,19 @@ "from torch.utils.data import Dataset, DataLoader\n", "from PIL import Image\n", "\n", + "\n", "class FilesDataset(Dataset):\n", - " \n", " def __init__(self, path, extension=\"*.jpg\"):\n", " self.path = Path(path)\n", " assert self.path.exists(), \"Path '{}' is not found\".format(path)\n", " self.images = list(self.path.rglob(extension))\n", " assert len(self.images) > 0, \"No images with extension {} found at '{}'\".format(extension, path)\n", - " \n", + "\n", " def __len__(self):\n", " return len(self.images)\n", - " \n", + "\n", " def __getitem__(self, i):\n", - " return Image.open(self.images[i]).convert('RGB')" + " return Image.open(self.images[i]).convert(\"RGB\")" ] }, { @@ -172,7 +173,7 @@ "train_A = FilesDataset(root / \"trainA\")\n", "train_B = FilesDataset(root / \"trainB\")\n", "\n", - "test_A = FilesDataset(root / \"testA\") \n", + "test_A = FilesDataset(root / \"testA\")\n", "test_B = FilesDataset(root / \"testB\")" ] }, @@ -189,7 +190,11 @@ "metadata": {}, "outputs": [], "source": [ - "print(\"Dataset sizes: \\ntrain A: {} | B: {}\\ntest A: {} | B: {}\\n\\t\".format(len(train_A), len(train_B), len(test_A), len(test_B)))" + "print(\n", + " \"Dataset sizes: \\ntrain A: {} | B: {}\\ntest A: {} | B: {}\\n\\t\".format(\n", + " len(train_A), len(train_B), len(test_A), len(test_B)\n", + " )\n", + ")" ] }, { @@ -198,7 +203,11 @@ "metadata": {}, "outputs": [], "source": [ - "print(\"Train random image sizes (A): {}, {}, {}, {}\".format(train_A[0].size, train_A[1].size, train_A[10].size, train_A[-1].size))" + "print(\n", + " \"Train random image sizes (A): {}, {}, {}, {}\".format(\n", + " train_A[0].size, train_A[1].size, train_A[10].size, train_A[-1].size\n", + " )\n", + ")" ] }, { @@ -207,7 +216,11 @@ "metadata": {}, "outputs": [], "source": [ - "print(\"Train random image sizes (B): {}, {}, {}, {}\".format(train_B[0].size, train_B[1].size, train_B[10].size, train_B[-1].size))" + "print(\n", + " \"Train random image sizes (B): {}, {}, {}, {}\".format(\n", + " train_B[0].size, train_B[1].size, train_B[10].size, train_B[-1].size\n", + " )\n", + ")" ] }, { @@ -217,6 +230,7 @@ "outputs": [], "source": [ "import matplotlib.pylab as plt\n", + "\n", "%matplotlib inline" ] }, @@ -267,11 +281,10 @@ "\n", "\n", "class Image2ImageDataset(Dataset):\n", - " \n", " def __init__(self, ds_a, ds_b):\n", " self.dataset_a = ds_a\n", " self.dataset_b = ds_b\n", - " \n", + "\n", " def __len__(self):\n", " return max(len(self.dataset_a), len(self.dataset_b))\n", "\n", @@ -279,21 +292,17 @@ " dp_a = self.dataset_a[i % len(self.dataset_a)]\n", " j = random.randint(0, len(self.dataset_b) - 1)\n", " dp_b = self.dataset_b[j]\n", - " return {\n", - " 'A': dp_a,\n", - " 'B': dp_b\n", - " }\n", + " return {\"A\": dp_a, \"B\": dp_b}\n", "\n", "\n", "class TransformedDataset(Dataset):\n", - " \n", " def __init__(self, ds, transform):\n", " self.dataset = ds\n", " self.transform = transform\n", - " \n", + "\n", " def __len__(self):\n", " return len(self.dataset)\n", - " \n", + "\n", " def __getitem__(self, i):\n", " return {k: self.transform(v) for k, v in self.dataset[i].items()}" ] @@ -319,10 +328,10 @@ "plt.figure(figsize=(10, 5))\n", "plt.subplot(121)\n", "plt.title(\"Train dataset 'Horses'\")\n", - "plt.imshow(dp['A'])\n", + "plt.imshow(dp[\"A\"])\n", "plt.subplot(122)\n", "plt.title(\"Train dataset 'Zebras'\")\n", - "plt.imshow(dp['B'])" + "plt.imshow(dp[\"B\"])" ] }, { @@ -336,10 +345,10 @@ "plt.figure(figsize=(10, 5))\n", "plt.subplot(121)\n", "plt.title(\"Test dataset 'Horses'\")\n", - "plt.imshow(dp['A'])\n", + "plt.imshow(dp[\"A\"])\n", "plt.subplot(122)\n", "plt.title(\"Test dataset 'Zebras'\")\n", - "plt.imshow(dp['B'])" + "plt.imshow(dp[\"B\"])" ] }, { @@ -351,27 +360,30 @@ "from torchvision.transforms import Compose, ColorJitter, RandomHorizontalFlip, ToTensor, Normalize, RandomCrop\n", "\n", "# To accelerate the training we reduce the image size to 200x200 pix instead of 256x256\n", - "train_transform = Compose([\n", - " RandomCrop(200),\n", - " RandomHorizontalFlip(),\n", - " ColorJitter(),\n", - " ToTensor(),\n", - " Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))\n", - "])\n", + "train_transform = Compose(\n", + " [\n", + " RandomCrop(200),\n", + " RandomHorizontalFlip(),\n", + " ColorJitter(),\n", + " ToTensor(),\n", + " Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5)),\n", + " ]\n", + ")\n", "transformed_train_ab_ds = TransformedDataset(train_ab_ds, transform=train_transform)\n", "\n", "# Please select appropriate batch_size value according to your infrastructure\n", "batch_size = 10\n", - "train_ab_loader = DataLoader(transformed_train_ab_ds, batch_size=batch_size, shuffle=True, drop_last=True, pin_memory=True, num_workers=4)\n", + "train_ab_loader = DataLoader(\n", + " transformed_train_ab_ds, batch_size=batch_size, shuffle=True, drop_last=True, pin_memory=True, num_workers=4\n", + ")\n", "\n", "\n", - "test_transform = Compose([\n", - " ToTensor(),\n", - " Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))\n", - "])\n", + "test_transform = Compose([ToTensor(), Normalize(mean=(0.5, 0.5, 0.5), std=(0.5, 0.5, 0.5))])\n", "transformed_test_ab_ds = TransformedDataset(test_ab_ds, transform=test_transform)\n", "batch_size = 10\n", - "test_ab_loader = DataLoader(transformed_test_ab_ds, batch_size=batch_size, shuffle=False, drop_last=False, pin_memory=True, num_workers=4)" + "test_ab_loader = DataLoader(\n", + " transformed_test_ab_ds, batch_size=batch_size, shuffle=False, drop_last=False, pin_memory=True, num_workers=4\n", + ")" ] }, { @@ -388,16 +400,12 @@ "plt.figure(figsize=(16, 8))\n", "plt.axis(\"off\")\n", "plt.title(\"Training Images from A\")\n", - "plt.imshow( \n", - " vutils.make_grid(real_batch['A'][:64], padding=2, normalize=True).cpu().numpy().transpose((1, 2, 0))\n", - ")\n", + "plt.imshow(vutils.make_grid(real_batch[\"A\"][:64], padding=2, normalize=True).cpu().numpy().transpose((1, 2, 0)))\n", "\n", "plt.figure(figsize=(16, 8))\n", "plt.axis(\"off\")\n", "plt.title(\"Training Images from B\")\n", - "plt.imshow(\n", - " vutils.make_grid(real_batch['B'][:64], padding=2, normalize=True).cpu().numpy().transpose((1, 2, 0))\n", - ")\n", + "plt.imshow(vutils.make_grid(real_batch[\"B\"][:64], padding=2, normalize=True).cpu().numpy().transpose((1, 2, 0)))\n", "real_batch = None\n", "torch.cuda.empty_cache()" ] @@ -445,29 +453,27 @@ " return nn.Sequential(\n", " nn.ConvTranspose2d(in_planes, out_planes, kernel_size=kernel_size, stride=stride, padding=1, output_padding=1),\n", " nn.InstanceNorm2d(out_planes, affine=False, track_running_stats=False),\n", - " nn.ReLU(inplace=True)\n", + " nn.ReLU(inplace=True),\n", " )\n", "\n", "\n", "class ResidualBlock(nn.Module):\n", - " \n", " def __init__(self, in_planes):\n", " super(ResidualBlock, self).__init__()\n", " self.conv1 = get_conv_inorm_relu(in_planes, in_planes, kernel_size=3, stride=1)\n", - " self.conv2 = get_conv_inorm_relu(in_planes, in_planes, kernel_size=3, stride=1, with_relu=False) \n", + " self.conv2 = get_conv_inorm_relu(in_planes, in_planes, kernel_size=3, stride=1, with_relu=False)\n", "\n", " def forward(self, x):\n", " residual = x\n", " x = self.conv1(x)\n", - " x = self.conv2(x) \n", + " x = self.conv2(x)\n", " return x + residual\n", "\n", "\n", "class Generator(nn.Module):\n", - " \n", " def __init__(self):\n", " super(Generator, self).__init__()\n", - " \n", + "\n", " self.c7s1_64 = get_conv_inorm_relu(3, 64, kernel_size=7, stride=1)\n", " self.d128 = get_conv_inorm_relu(64, 128, kernel_size=3, stride=2, reflection_pad=False)\n", " self.d256 = get_conv_inorm_relu(128, 256, kernel_size=3, stride=2, reflection_pad=False)\n", @@ -485,7 +491,7 @@ " x = self.c7s1_64(x)\n", " x = self.d128(x)\n", " x = self.d256(x)\n", - " \n", + "\n", " # 9 residual blocks\n", " x = self.resnet9(x)\n", "\n", @@ -493,7 +499,7 @@ " x = self.u128(x)\n", " x = self.u64(x)\n", " y = self.c7s1_3(x)\n", - " return y\n" + " return y" ] }, { @@ -548,12 +554,11 @@ " return nn.Sequential(\n", " nn.Conv2d(in_planes, out_planes, kernel_size=4, stride=stride, padding=1),\n", " nn.InstanceNorm2d(out_planes, affine=False, track_running_stats=False),\n", - " nn.LeakyReLU(negative_slope=negative_slope, inplace=True)\n", + " nn.LeakyReLU(negative_slope=negative_slope, inplace=True),\n", " )\n", "\n", "\n", "class Discriminator(nn.Module):\n", - "\n", " def __init__(self):\n", " super(Discriminator, self).__init__()\n", " self.c64 = nn.Conv2d(3, 64, kernel_size=4, stride=2, padding=1)\n", @@ -571,7 +576,7 @@ " x = self.c256(x)\n", " x = self.c512(x)\n", " y = self.last_conv(x)\n", - " return y\n" + " return y" ] }, { @@ -622,7 +627,8 @@ "metadata": {}, "outputs": [], "source": [ - "g = None; d = None" + "g = None\n", + "d = None" ] }, { @@ -735,7 +741,7 @@ " buffer.append(b.cpu())\n", " elif random.uniform(0, 1) > 0.5:\n", " # Add newly created image into the buffer and put ont from the buffer into the output\n", - " random_index = random.randint(0, buffer_size - 1) \n", + " random_index = random.randint(0, buffer_size - 1)\n", " output_batch.append(buffer[random_index].clone().to(device))\n", " buffer[random_index] = b.cpu()\n", " else:\n", @@ -783,7 +789,7 @@ "\n", "def discriminator_forward_pass(discriminator, batch_real, batch_fake, fake_buffer):\n", " decision_real = discriminator(batch_real)\n", - " batch_fake = buffer_insert_and_get(fake_buffer, batch_fake) \n", + " batch_fake = buffer_insert_and_get(fake_buffer, batch_fake)\n", " decision_fake = discriminator(batch_fake)\n", " return decision_real, decision_fake\n", "\n", @@ -793,12 +799,12 @@ " target = torch.ones_like(batch_decision)\n", " loss_gan = F.mse_loss(batch_decision, target)\n", " # loss cycle\n", - " loss_cycle = F.l1_loss(batch_rec, batch_real) * lambda_value \n", + " loss_cycle = F.l1_loss(batch_rec, batch_real) * lambda_value\n", " return loss_gan + loss_cycle\n", "\n", "\n", "def compute_loss_discriminator(decision_real, decision_fake):\n", - " # loss = mean (D_b(y) − 1)^2 + mean D_b(G(x))^2 \n", + " # loss = mean (D_b(y) − 1)^2 + mean D_b(G(x))^2\n", " loss = F.mse_loss(decision_fake, torch.zeros_like(decision_fake))\n", " loss += F.mse_loss(decision_real, torch.ones_like(decision_real))\n", " return loss\n", @@ -808,16 +814,16 @@ " generator_A2B.train()\n", " generator_B2A.train()\n", " discriminator_A.train()\n", - " discriminator_B.train() \n", + " discriminator_B.train()\n", "\n", - " real_a = convert_tensor(batch['A'], device=device, non_blocking=True)\n", - " real_b = convert_tensor(batch['B'], device=device, non_blocking=True)\n", + " real_a = convert_tensor(batch[\"A\"], device=device, non_blocking=True)\n", + " real_b = convert_tensor(batch[\"B\"], device=device, non_blocking=True)\n", "\n", " # Update generators\n", "\n", " # Disable grads computation for the discriminators:\n", " toggle_grad(discriminator_A, False)\n", - " toggle_grad(discriminator_B, False) \n", + " toggle_grad(discriminator_B, False)\n", "\n", " with autocast(enabled=amp_enabled):\n", " fake_b = generator_A2B(real_a)\n", @@ -828,8 +834,8 @@ " decision_fake_b = discriminator_B(fake_b)\n", "\n", " # Compute loss for generators and update generators\n", - " # loss_a2b = GAN loss: mean (D_b(G(x)) − 1)^2 + Forward cycle loss: || F(G(x)) - x ||_1 \n", - " loss_a2b = compute_loss_generator(decision_fake_b, real_a, rec_a, lambda_value) \n", + " # loss_a2b = GAN loss: mean (D_b(G(x)) − 1)^2 + Forward cycle loss: || F(G(x)) - x ||_1\n", + " loss_a2b = compute_loss_generator(decision_fake_b, real_a, rec_a, lambda_value)\n", "\n", " # loss_b2a = GAN loss: mean (D_a(F(x)) − 1)^2 + Backward cycle loss: || G(F(y)) - y ||_1\n", " loss_b2a = compute_loss_generator(decision_fake_a, real_b, rec_b, lambda_value)\n", @@ -842,29 +848,33 @@ " amp_scaler.step(optimizer_G)\n", "\n", " decision_fake_a = rec_a = decision_fake_b = rec_b = None\n", - " \n", + "\n", " # Enable grads computation for the discriminators:\n", " toggle_grad(discriminator_A, True)\n", - " toggle_grad(discriminator_B, True) \n", + " toggle_grad(discriminator_B, True)\n", "\n", " with autocast(enabled=amp_enabled):\n", - " decision_real_a, decision_fake_a = discriminator_forward_pass(discriminator_A, real_a, fake_a.detach(), fake_a_buffer) \n", - " decision_real_b, decision_fake_b = discriminator_forward_pass(discriminator_B, real_b, fake_b.detach(), fake_b_buffer) \n", + " decision_real_a, decision_fake_a = discriminator_forward_pass(\n", + " discriminator_A, real_a, fake_a.detach(), fake_a_buffer\n", + " )\n", + " decision_real_b, decision_fake_b = discriminator_forward_pass(\n", + " discriminator_B, real_b, fake_b.detach(), fake_b_buffer\n", + " )\n", " # Compute loss for discriminators and update discriminators\n", " # loss_a = mean (D_a(y) − 1)^2 + mean D_a(F(x))^2\n", " loss_a = compute_loss_discriminator(decision_real_a, decision_fake_a)\n", "\n", " # loss_b = mean (D_b(y) − 1)^2 + mean D_b(G(x))^2\n", " loss_b = compute_loss_discriminator(decision_real_b, decision_fake_b)\n", - " \n", + "\n", " # total discriminators loss:\n", " loss_discriminators = 0.5 * (loss_a + loss_b)\n", - " \n", + "\n", " optimizer_D.zero_grad()\n", " amp_scaler.scale(loss_discriminators).backward()\n", " amp_scaler.step(optimizer_D)\n", " amp_scaler.update()\n", - " \n", + "\n", " return {\n", " \"loss_generators\": loss_generators.item(),\n", " \"loss_generator_a2b\": loss_a2b.item(),\n", @@ -872,7 +882,7 @@ " \"loss_discriminators\": loss_discriminators.item(),\n", " \"loss_discriminator_a\": loss_a.item(),\n", " \"loss_discriminator_b\": loss_b.item(),\n", - " }\n" + " }" ] }, { @@ -949,20 +959,22 @@ "trainer = Engine(update_fn)\n", "\n", "metric_names = [\n", - " 'loss_discriminators', \n", - " 'loss_generators', \n", - " 'loss_discriminator_a',\n", - " 'loss_discriminator_b',\n", - " 'loss_generator_a2b',\n", - " 'loss_generator_b2a' \n", + " \"loss_discriminators\",\n", + " \"loss_generators\",\n", + " \"loss_discriminator_a\",\n", + " \"loss_discriminator_b\",\n", + " \"loss_generator_a2b\",\n", + " \"loss_generator_b2a\",\n", "]\n", "\n", + "\n", "def output_transform(out, name):\n", " return out[name]\n", "\n", + "\n", "for name in metric_names:\n", " # here we cannot use lambdas as they do not store argument `name`\n", - " RunningAverage(output_transform=partial(output_transform, name=name)).attach(trainer, name)\n" + " RunningAverage(output_transform=partial(output_transform, name=name)).attach(trainer, name)" ] }, { @@ -976,9 +988,7 @@ "exp_name = datetime.now().strftime(\"%Y%m%d-%H%M%S\")\n", "tb_logger = TensorboardLogger(log_dir=\"/tmp/cycle_gan_horse2zebra_tb_logs/{}\".format(exp_name))\n", "\n", - "tb_logger.attach(trainer, \n", - " log_handler=OutputHandler('training', metric_names), \n", - " event_name=Events.ITERATION_COMPLETED)\n", + "tb_logger.attach(trainer, log_handler=OutputHandler(\"training\", metric_names), event_name=Events.ITERATION_COMPLETED)\n", "\n", "print(\"Experiment name: \", exp_name)" ] @@ -997,11 +1007,7 @@ " if not wb_dir.exists():\n", " wb_dir.mkdir()\n", " wb_logger = WandBLogger(\n", - " project=\"ignite-cyclegan-torch-amp\",\n", - " name=wb_run_name,\n", - " sync_tensorboard=True,\n", - " dir=wb_dir.as_posix(),\n", - " reinit=True\n", + " project=\"ignite-cyclegan-torch-amp\", name=wb_run_name, sync_tensorboard=True, dir=wb_dir.as_posix(), reinit=True\n", " )\n", "except RuntimeError:\n", " wb_logger = None" @@ -1025,24 +1031,24 @@ "\n", "def evaluate_fn(engine, batch):\n", " generator_A2B.eval()\n", - " generator_B2A.eval() \n", + " generator_B2A.eval()\n", " with torch.no_grad():\n", - " real_a = convert_tensor(batch['A'], device=device, non_blocking=True)\n", - " real_b = convert_tensor(batch['B'], device=device, non_blocking=True)\n", - " \n", + " real_a = convert_tensor(batch[\"A\"], device=device, non_blocking=True)\n", + " real_b = convert_tensor(batch[\"B\"], device=device, non_blocking=True)\n", + "\n", " fake_b = generator_A2B(real_a)\n", " rec_a = generator_B2A(fake_b)\n", "\n", " fake_a = generator_B2A(real_b)\n", " rec_b = generator_A2B(fake_a)\n", - " \n", + "\n", " return {\n", - " 'real_a': real_a,\n", - " 'real_b': real_b,\n", - " 'fake_a': fake_a,\n", - " 'fake_b': fake_b,\n", - " 'rec_a': rec_a,\n", - " 'rec_b': rec_b, \n", + " \"real_a\": real_a,\n", + " \"real_b\": real_b,\n", + " \"fake_a\": fake_a,\n", + " \"fake_b\": fake_b,\n", + " \"rec_a\": rec_a,\n", + " \"rec_b\": rec_b,\n", " }\n", "\n", "\n", @@ -1067,8 +1073,12 @@ "small_test_ds = Subset(test_ab_ds, test_random_indices)\n", "small_test_ds = TransformedDataset(small_test_ds, transform=test_transform)\n", "\n", - "eval_train_loader = DataLoader(small_train_ds, batch_size=eval_batch_size, shuffle=False, drop_last=False, pin_memory=True, num_workers=4)\n", - "eval_test_loader = DataLoader(small_test_ds, batch_size=eval_batch_size, shuffle=False, drop_last=False, pin_memory=True, num_workers=4)" + "eval_train_loader = DataLoader(\n", + " small_train_ds, batch_size=eval_batch_size, shuffle=False, drop_last=False, pin_memory=True, num_workers=4\n", + ")\n", + "eval_test_loader = DataLoader(\n", + " small_test_ds, batch_size=eval_batch_size, shuffle=False, drop_last=False, pin_memory=True, num_workers=4\n", + ")" ] }, { @@ -1084,7 +1094,6 @@ "\n", "\n", "def log_generated_images(engine, logger, event_name):\n", - "\n", " tag = \"Train\" if engine.state.dataloader == eval_train_loader else \"Test\"\n", " output = engine.state.output\n", " state = trainer.state\n", @@ -1094,30 +1103,50 @@ " # [real a1, real a2, ...]\n", " # [fake a1, fake a2, ...]\n", " # [rec a1, rec a2, ...]\n", - " \n", - " s = output['real_a'].shape[0]\n", - " res_a = vutils.make_grid(torch.cat([\n", - " output['real_a'],\n", - " output['fake_b'],\n", - " output['rec_a'],\n", - " ]), padding=2, normalize=True, nrow=s).cpu()\n", "\n", - " logger.writer.add_image(tag=\"{} Horses2Zebras (real, fake, rec)\".format(tag), \n", - " img_tensor=res_a, global_step=global_step, dataformats='CHW')\n", + " s = output[\"real_a\"].shape[0]\n", + " res_a = vutils.make_grid(\n", + " torch.cat(\n", + " [\n", + " output[\"real_a\"],\n", + " output[\"fake_b\"],\n", + " output[\"rec_a\"],\n", + " ]\n", + " ),\n", + " padding=2,\n", + " normalize=True,\n", + " nrow=s,\n", + " ).cpu()\n", + "\n", + " logger.writer.add_image(\n", + " tag=\"{} Horses2Zebras (real, fake, rec)\".format(tag),\n", + " img_tensor=res_a,\n", + " global_step=global_step,\n", + " dataformats=\"CHW\",\n", + " )\n", + "\n", + " s = output[\"real_b\"].shape[0]\n", + " res_b = vutils.make_grid(\n", + " torch.cat(\n", + " [\n", + " output[\"real_b\"],\n", + " output[\"fake_a\"],\n", + " output[\"rec_b\"],\n", + " ]\n", + " ),\n", + " padding=2,\n", + " normalize=True,\n", + " nrow=s,\n", + " ).cpu()\n", + " logger.writer.add_image(\n", + " tag=\"{} Zebras2Horses (real, fake, rec)\".format(tag),\n", + " img_tensor=res_b,\n", + " global_step=global_step,\n", + " dataformats=\"CHW\",\n", + " )\n", "\n", - " s = output['real_b'].shape[0]\n", - " res_b = vutils.make_grid(torch.cat([\n", - " output['real_b'],\n", - " output['fake_a'],\n", - " output['rec_b'],\n", - " ]), padding=2, normalize=True, nrow=s).cpu()\n", - " logger.writer.add_image(tag=\"{} Zebras2Horses (real, fake, rec)\".format(tag), \n", - " img_tensor=res_b, global_step=global_step, dataformats='CHW')\n", "\n", - " \n", - "tb_logger.attach(evaluator,\n", - " log_handler=log_generated_images, \n", - " event_name=Events.COMPLETED)" + "tb_logger.attach(evaluator, log_handler=log_generated_images, event_name=Events.COMPLETED)" ] }, { @@ -1137,23 +1166,18 @@ "\n", "lr = 0.0002\n", "\n", - "milestones_values = [\n", - " (0, lr),\n", - " (100, lr),\n", - " (200, 0.0)\n", - "]\n", - "gen_lr_scheduler = PiecewiseLinear(optimizer_D, param_name='lr', milestones_values=milestones_values)\n", - "desc_lr_scheduler = PiecewiseLinear(optimizer_G, param_name='lr', milestones_values=milestones_values)\n", + "milestones_values = [(0, lr), (100, lr), (200, 0.0)]\n", + "gen_lr_scheduler = PiecewiseLinear(optimizer_D, param_name=\"lr\", milestones_values=milestones_values)\n", + "desc_lr_scheduler = PiecewiseLinear(optimizer_G, param_name=\"lr\", milestones_values=milestones_values)\n", "\n", - "lr_scheduler = ParamGroupScheduler([gen_lr_scheduler, desc_lr_scheduler], \n", - " names=['gen_lr_scheduler', 'desc_lr_scheduler'])\n", + "lr_scheduler = ParamGroupScheduler(\n", + " [gen_lr_scheduler, desc_lr_scheduler], names=[\"gen_lr_scheduler\", \"desc_lr_scheduler\"]\n", + ")\n", "\n", "trainer.add_event_handler(Events.EPOCH_STARTED, lr_scheduler)\n", "\n", "\n", - "tb_logger.attach(trainer,\n", - " log_handler=OptimizerParamsHandler(optimizer_G, \"lr\"), \n", - " event_name=Events.EPOCH_STARTED)" + "tb_logger.attach(trainer, log_handler=OptimizerParamsHandler(optimizer_G, \"lr\"), event_name=Events.EPOCH_STARTED)" ] }, { @@ -1191,7 +1215,6 @@ " \"discriminator_B\": discriminator_B,\n", " \"generator_B2A\": generator_B2A,\n", " \"discriminator_A\": discriminator_A,\n", - " \n", " \"optimizer_G\": optimizer_G,\n", " \"optimizer_D\": optimizer_D,\n", "}\n", @@ -1211,8 +1234,12 @@ "# Iteration-wise progress bar\n", "ProgressBar(bar_format=\"\").attach(trainer)\n", "# Epoch-wise progress bar with display of training losses\n", - "ProgressBar(persist=True, bar_format=\"\").attach(trainer, metric_names=['loss_discriminators', 'loss_generators'], \n", - " event_name=Events.EPOCH_STARTED, closing_event_name=Events.COMPLETED)" + "ProgressBar(persist=True, bar_format=\"\").attach(\n", + " trainer,\n", + " metric_names=[\"loss_discriminators\", \"loss_generators\"],\n", + " event_name=Events.EPOCH_STARTED,\n", + " closing_event_name=Events.COMPLETED,\n", + ")" ] }, { @@ -1307,16 +1334,16 @@ "outputs": [], "source": [ "i = random.randint(0, len(test_ab_ds) - 1)\n", - "img = test_ab_ds[i]['A']\n", + "img = test_ab_ds[i][\"A\"]\n", "x = test_transform(img)\n", "x = x.unsqueeze(0).to(device)\n", "\n", "\n", "with torch.no_grad():\n", " y_pred = generator_A2B(x)\n", - " \n", "\n", - "img_pred = (255 * normalize(y_pred[0, ...])).cpu().numpy().transpose((1, 2, 0)).astype('uint8')" + "\n", + "img_pred = (255 * normalize(y_pred[0, ...])).cpu().numpy().transpose((1, 2, 0)).astype(\"uint8\")" ] }, { @@ -1367,7 +1394,7 @@ " y_pred = generator_A2B(x)\n", "\n", "\n", - "img_pred = (255 * normalize(y_pred[0, ...])).cpu().numpy().transpose((1, 2, 0)).astype('uint8')" + "img_pred = (255 * normalize(y_pred[0, ...])).cpu().numpy().transpose((1, 2, 0)).astype(\"uint8\")" ] }, { diff --git a/examples/notebooks/EfficientNet_Cifar100_finetuning.ipynb b/examples/notebooks/EfficientNet_Cifar100_finetuning.ipynb index 040cd627e2e6..413b06740361 100644 --- a/examples/notebooks/EfficientNet_Cifar100_finetuning.ipynb +++ b/examples/notebooks/EfficientNet_Cifar100_finetuning.ipynb @@ -148,15 +148,13 @@ "\n", "\n", "class Swish(nn.Module):\n", - " \n", " def forward(self, x):\n", " return x * torch.sigmoid(x)\n", "\n", "\n", "class Flatten(nn.Module):\n", - " \n", " def forward(self, x):\n", - " return x.reshape(x.shape[0], -1)\n" + " return x.reshape(x.shape[0], -1)" ] }, { @@ -173,6 +171,7 @@ "outputs": [], "source": [ "import matplotlib.pylab as plt\n", + "\n", "%matplotlib inline\n", "\n", "d = torch.linspace(-10.0, 10.0)\n", @@ -181,8 +180,8 @@ "res2 = torch.relu(d)\n", "\n", "plt.title(\"Swish transformation\")\n", - "plt.plot(d.numpy(), res.numpy(), label='Swish')\n", - "plt.plot(d.numpy(), res2.numpy(), label='ReLU')\n", + "plt.plot(d.numpy(), res.numpy(), label=\"Swish\")\n", + "plt.plot(d.numpy(), res2.numpy(), label=\"ReLU\")\n", "plt.legend()" ] }, @@ -200,22 +199,19 @@ "outputs": [], "source": [ "class SqueezeExcitation(nn.Module):\n", - " \n", " def __init__(self, inplanes, se_planes):\n", " super(SqueezeExcitation, self).__init__()\n", " self.reduce_expand = nn.Sequential(\n", - " nn.Conv2d(inplanes, se_planes, \n", - " kernel_size=1, stride=1, padding=0, bias=True),\n", + " nn.Conv2d(inplanes, se_planes, kernel_size=1, stride=1, padding=0, bias=True),\n", " Swish(),\n", - " nn.Conv2d(se_planes, inplanes, \n", - " kernel_size=1, stride=1, padding=0, bias=True),\n", - " nn.Sigmoid()\n", + " nn.Conv2d(se_planes, inplanes, kernel_size=1, stride=1, padding=0, bias=True),\n", + " nn.Sigmoid(),\n", " )\n", "\n", " def forward(self, x):\n", " x_se = torch.mean(x, dim=(-2, -1), keepdim=True)\n", " x_se = self.reduce_expand(x_se)\n", - " return x_se * x\n" + " return x_se * x" ] }, { @@ -238,52 +234,52 @@ "\n", "\n", "class MBConv(nn.Module):\n", - "\n", - " def __init__(self, inplanes, planes, kernel_size, stride, \n", - " expand_rate=1.0, se_rate=0.25, \n", - " drop_connect_rate=0.2):\n", + " def __init__(self, inplanes, planes, kernel_size, stride, expand_rate=1.0, se_rate=0.25, drop_connect_rate=0.2):\n", " super(MBConv, self).__init__()\n", "\n", " expand_planes = int(inplanes * expand_rate)\n", " se_planes = max(1, int(inplanes * se_rate))\n", "\n", - " self.expansion_conv = None \n", + " self.expansion_conv = None\n", " if expand_rate > 1.0:\n", " self.expansion_conv = nn.Sequential(\n", - " nn.Conv2d(inplanes, expand_planes, \n", - " kernel_size=1, stride=1, padding=0, bias=False),\n", + " nn.Conv2d(inplanes, expand_planes, kernel_size=1, stride=1, padding=0, bias=False),\n", " nn.BatchNorm2d(expand_planes, momentum=0.01, eps=1e-3),\n", - " Swish()\n", + " Swish(),\n", " )\n", " inplanes = expand_planes\n", "\n", " self.depthwise_conv = nn.Sequential(\n", - " nn.Conv2d(inplanes, expand_planes,\n", - " kernel_size=kernel_size, stride=stride, \n", - " padding=kernel_size // 2, groups=expand_planes,\n", - " bias=False),\n", + " nn.Conv2d(\n", + " inplanes,\n", + " expand_planes,\n", + " kernel_size=kernel_size,\n", + " stride=stride,\n", + " padding=kernel_size // 2,\n", + " groups=expand_planes,\n", + " bias=False,\n", + " ),\n", " nn.BatchNorm2d(expand_planes, momentum=0.01, eps=1e-3),\n", - " Swish()\n", + " Swish(),\n", " )\n", "\n", " self.squeeze_excitation = SqueezeExcitation(expand_planes, se_planes)\n", - " \n", + "\n", " self.project_conv = nn.Sequential(\n", - " nn.Conv2d(expand_planes, planes, \n", - " kernel_size=1, stride=1, padding=0, bias=False),\n", + " nn.Conv2d(expand_planes, planes, kernel_size=1, stride=1, padding=0, bias=False),\n", " nn.BatchNorm2d(planes, momentum=0.01, eps=1e-3),\n", " )\n", "\n", " self.with_skip = stride == 1\n", " self.drop_connect_rate = drop_connect_rate\n", - " \n", - " def _drop_connect(self, x): \n", + "\n", + " def _drop_connect(self, x):\n", " keep_prob = 1.0 - self.drop_connect_rate\n", " drop_mask = torch.rand(x.shape[0], 1, 1, 1) + keep_prob\n", " drop_mask = drop_mask.type_as(x)\n", " drop_mask.floor_()\n", " return drop_mask * x / keep_prob\n", - " \n", + "\n", " def forward(self, x):\n", " z = x\n", " if self.expansion_conv is not None:\n", @@ -292,9 +288,9 @@ " x = self.depthwise_conv(x)\n", " x = self.squeeze_excitation(x)\n", " x = self.project_conv(x)\n", - " \n", + "\n", " # Add identity skip\n", - " if x.shape == z.shape and self.with_skip: \n", + " if x.shape == z.shape and self.with_skip:\n", " if self.training and self.drop_connect_rate is not None:\n", " x = self._drop_connect(x)\n", " x += z\n", @@ -318,19 +314,18 @@ "import math\n", "\n", "\n", - "def init_weights(module): \n", - " if isinstance(module, nn.Conv2d): \n", - " nn.init.kaiming_normal_(module.weight, a=0, mode='fan_out')\n", + "def init_weights(module):\n", + " if isinstance(module, nn.Conv2d):\n", + " nn.init.kaiming_normal_(module.weight, a=0, mode=\"fan_out\")\n", " elif isinstance(module, nn.Linear):\n", " init_range = 1.0 / math.sqrt(module.weight.shape[1])\n", " nn.init.uniform_(module.weight, a=-init_range, b=init_range)\n", - " \n", - " \n", + "\n", + "\n", "class EfficientNet(nn.Module):\n", - " \n", " def _setup_repeats(self, num_repeats):\n", " return int(math.ceil(self.depth_coefficient * num_repeats))\n", - " \n", + "\n", " def _setup_channels(self, num_channels):\n", " num_channels *= self.width_coefficient\n", " new_num_channels = math.floor(num_channels / self.divisor + 0.5) * self.divisor\n", @@ -339,24 +334,27 @@ " new_num_channels += self.divisor\n", " return new_num_channels\n", "\n", - " def __init__(self, num_classes=100, \n", - " width_coefficient=1.0,\n", - " depth_coefficient=1.0,\n", - " se_rate=0.25,\n", - " dropout_rate=0.2,\n", - " drop_connect_rate=0.2):\n", + " def __init__(\n", + " self,\n", + " num_classes=100,\n", + " width_coefficient=1.0,\n", + " depth_coefficient=1.0,\n", + " se_rate=0.25,\n", + " dropout_rate=0.2,\n", + " drop_connect_rate=0.2,\n", + " ):\n", " super(EfficientNet, self).__init__()\n", - " \n", + "\n", " self.width_coefficient = width_coefficient\n", " self.depth_coefficient = depth_coefficient\n", " self.divisor = 8\n", - " \n", + "\n", " list_channels = [32, 16, 24, 40, 80, 112, 192, 320, 1280]\n", " list_channels = [self._setup_channels(c) for c in list_channels]\n", - " \n", + "\n", " list_num_repeats = [1, 2, 2, 3, 3, 4, 1]\n", - " list_num_repeats = [self._setup_repeats(r) for r in list_num_repeats] \n", - " \n", + " list_num_repeats = [self._setup_repeats(r) for r in list_num_repeats]\n", + "\n", " expand_rates = [1, 6, 6, 6, 6, 6, 6]\n", " strides = [1, 2, 2, 2, 1, 2, 1]\n", " kernel_sizes = [3, 3, 5, 3, 5, 5, 3]\n", @@ -365,15 +363,14 @@ " self.stem = nn.Sequential(\n", " nn.Conv2d(3, list_channels[0], kernel_size=3, stride=2, padding=1, bias=False),\n", " nn.BatchNorm2d(list_channels[0], momentum=0.01, eps=1e-3),\n", - " Swish()\n", + " Swish(),\n", " )\n", - " \n", + "\n", " # Define MBConv blocks\n", " blocks = []\n", " counter = 0\n", " num_blocks = sum(list_num_repeats)\n", " for idx in range(7):\n", - " \n", " num_channels = list_channels[idx]\n", " next_num_channels = list_channels[idx + 1]\n", " num_repeats = list_num_repeats[idx]\n", @@ -381,42 +378,57 @@ " kernel_size = kernel_sizes[idx]\n", " stride = strides[idx]\n", " drop_rate = drop_connect_rate * counter / num_blocks\n", - " \n", + "\n", " name = \"MBConv{}_{}\".format(expand_rate, counter)\n", - " blocks.append((\n", - " name,\n", - " MBConv(num_channels, next_num_channels, \n", - " kernel_size=kernel_size, stride=stride, expand_rate=expand_rate, \n", - " se_rate=se_rate, drop_connect_rate=drop_rate)\n", - " ))\n", + " blocks.append(\n", + " (\n", + " name,\n", + " MBConv(\n", + " num_channels,\n", + " next_num_channels,\n", + " kernel_size=kernel_size,\n", + " stride=stride,\n", + " expand_rate=expand_rate,\n", + " se_rate=se_rate,\n", + " drop_connect_rate=drop_rate,\n", + " ),\n", + " )\n", + " )\n", " counter += 1\n", - " for i in range(1, num_repeats): \n", + " for i in range(1, num_repeats):\n", " name = \"MBConv{}_{}\".format(expand_rate, counter)\n", - " drop_rate = drop_connect_rate * counter / num_blocks \n", - " blocks.append((\n", - " name,\n", - " MBConv(next_num_channels, next_num_channels, \n", - " kernel_size=kernel_size, stride=1, expand_rate=expand_rate, \n", - " se_rate=se_rate, drop_connect_rate=drop_rate) \n", - " ))\n", + " drop_rate = drop_connect_rate * counter / num_blocks\n", + " blocks.append(\n", + " (\n", + " name,\n", + " MBConv(\n", + " next_num_channels,\n", + " next_num_channels,\n", + " kernel_size=kernel_size,\n", + " stride=1,\n", + " expand_rate=expand_rate,\n", + " se_rate=se_rate,\n", + " drop_connect_rate=drop_rate,\n", + " ),\n", + " )\n", + " )\n", " counter += 1\n", - " \n", + "\n", " self.blocks = nn.Sequential(OrderedDict(blocks))\n", - " \n", + "\n", " # Define head\n", " self.head = nn.Sequential(\n", - " nn.Conv2d(list_channels[-2], list_channels[-1], \n", - " kernel_size=1, bias=False),\n", + " nn.Conv2d(list_channels[-2], list_channels[-1], kernel_size=1, bias=False),\n", " nn.BatchNorm2d(list_channels[-1], momentum=0.01, eps=1e-3),\n", " Swish(),\n", " nn.AdaptiveAvgPool2d(1),\n", " Flatten(),\n", " nn.Dropout(p=dropout_rate),\n", - " nn.Linear(list_channels[-1], num_classes)\n", + " nn.Linear(list_channels[-1], num_classes),\n", " )\n", "\n", " self.apply(init_weights)\n", - " \n", + "\n", " def forward(self, x):\n", " f = self.stem(x)\n", " f = self.blocks(f)\n", @@ -449,9 +461,7 @@ "metadata": {}, "outputs": [], "source": [ - "model = EfficientNet(num_classes=1000, \n", - " width_coefficient=1.0, depth_coefficient=1.0, \n", - " dropout_rate=0.2)" + "model = EfficientNet(num_classes=1000, width_coefficient=1.0, depth_coefficient=1.0, dropout_rate=0.2)" ] }, { @@ -473,11 +483,12 @@ " num_params = 1\n", " for s in p.shape:\n", " num_params *= s\n", - " if display_all_modules: print(\"{}: {}\".format(n, num_params))\n", + " if display_all_modules:\n", + " print(\"{}: {}\".format(n, num_params))\n", " total_num_params += num_params\n", " print(\"-\" * 50)\n", " print(\"Total number of parameters: {:.2e}\".format(total_num_params))\n", - " \n", + "\n", "\n", "print_num_params(model)" ] @@ -534,7 +545,7 @@ "\n", "def show_graph(graph_def):\n", " \"\"\"Visualize TensorFlow graph.\"\"\"\n", - " if hasattr(graph_def, 'as_graph_def'):\n", + " if hasattr(graph_def, \"as_graph_def\"):\n", " graph_def = graph_def.as_graph_def()\n", " strip_def = graph_def\n", " code = \"\"\"\n", @@ -548,11 +559,11 @@ "
\n", " \n", "
\n", - " \"\"\".format(data=repr(str(strip_def)), id='graph'+str(random.randint(0, 1000)))\n", + " \"\"\".format(data=repr(str(strip_def)), id=\"graph\" + str(random.randint(0, 1000)))\n", "\n", " iframe = \"\"\"\n", " \n", - " \"\"\".format(code.replace('\"', '"'))\n", + " \"\"\".format(code.replace('\"', \""\"))\n", " display(HTML(iframe))" ] }, @@ -565,7 +576,7 @@ "x = torch.rand(4, 3, 224, 224)\n", "\n", "# Error : module 'torch.onnx' has no attribute 'set_training'\n", - "# uncomment when it will be fixed \n", + "# uncomment when it will be fixed\n", "\n", "# graph_def = graph(model, x, operator_export_type='RAW')" ] @@ -610,12 +621,8 @@ "model_state = torch.load(\"/tmp/efficientnet_weights/efficientnet-b0-08094119.pth\")\n", "\n", "# A basic remapping is required\n", - "mapping = {\n", - " k: v for k, v in zip(model_state.keys(), model.state_dict().keys())\n", - "}\n", - "mapped_model_state = OrderedDict([\n", - " (mapping[k], v) for k, v in model_state.items()\n", - "])\n", + "mapping = {k: v for k, v in zip(model_state.keys(), model.state_dict().keys())}\n", + "mapped_model_state = OrderedDict([(mapping[k], v) for k, v in model_state.items()])\n", "\n", "model.load_state_dict(mapped_model_state, strict=False)" ] @@ -648,10 +655,14 @@ "img = Image.open(\"/tmp/giant_panda.jpg\")\n", "# Preprocess image\n", "image_size = 224\n", - "tfms = transforms.Compose([transforms.Resize(image_size), \n", - " transforms.CenterCrop(image_size), \n", - " transforms.ToTensor(),\n", - " transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),])\n", + "tfms = transforms.Compose(\n", + " [\n", + " transforms.Resize(image_size),\n", + " transforms.CenterCrop(image_size),\n", + " transforms.ToTensor(),\n", + " transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]),\n", + " ]\n", + ")\n", "x = tfms(img).unsqueeze(0)\n", "\n", "plt.imshow(img)" @@ -669,10 +680,10 @@ " y_pred = model(x)\n", "\n", "# Print predictions\n", - "print('-----')\n", + "print(\"-----\")\n", "for idx in torch.topk(y_pred, k=5).indices.squeeze(0).tolist():\n", " prob = torch.softmax(y_pred, dim=1)[0, idx].item()\n", - " print('{label:<75} ({p:.2f}%)'.format(label=labels[str(idx)], p=prob*100))" + " print(\"{label:<75} ({p:.2f}%)\".format(label=labels[str(idx)], p=prob * 100))" ] }, { @@ -699,7 +710,7 @@ "metadata": {}, "outputs": [], "source": [ - "from torchvision.datasets.cifar import CIFAR100 \n", + "from torchvision.datasets.cifar import CIFAR100\n", "from torchvision.transforms import Compose, RandomCrop, Pad, RandomHorizontalFlip, Resize\n", "from torchvision.transforms import ToTensor, Normalize\n", "\n", @@ -717,19 +728,19 @@ "from PIL.Image import BICUBIC\n", "\n", "\n", - "train_transform = Compose([\n", - " Resize(256, BICUBIC),\n", - " RandomCrop(224),\n", - " RandomHorizontalFlip(),\n", - " ToTensor(),\n", - " Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n", - "])\n", + "train_transform = Compose(\n", + " [\n", + " Resize(256, BICUBIC),\n", + " RandomCrop(224),\n", + " RandomHorizontalFlip(),\n", + " ToTensor(),\n", + " Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n", + " ]\n", + ")\n", "\n", - "test_transform = Compose([\n", - " Resize(224, BICUBIC), \n", - " ToTensor(),\n", - " Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n", - "])\n", + "test_transform = Compose(\n", + " [Resize(224, BICUBIC), ToTensor(), Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])]\n", + ")\n", "\n", "\n", "train_dataset = CIFAR100(root=path, train=True, transform=train_transform, download=True)\n", @@ -753,14 +764,17 @@ "\n", "batch_size = 172\n", "\n", - "train_loader = DataLoader(train_dataset, batch_size=batch_size, num_workers=20, \n", - " shuffle=True, drop_last=True, pin_memory=True)\n", + "train_loader = DataLoader(\n", + " train_dataset, batch_size=batch_size, num_workers=20, shuffle=True, drop_last=True, pin_memory=True\n", + ")\n", "\n", - "test_loader = DataLoader(test_dataset, batch_size=batch_size, num_workers=20, \n", - " shuffle=False, drop_last=False, pin_memory=True)\n", + "test_loader = DataLoader(\n", + " test_dataset, batch_size=batch_size, num_workers=20, shuffle=False, drop_last=False, pin_memory=True\n", + ")\n", "\n", - "eval_train_loader = DataLoader(train_eval_dataset, batch_size=batch_size, num_workers=20, \n", - " shuffle=False, drop_last=False, pin_memory=True)" + "eval_train_loader = DataLoader(\n", + " train_eval_dataset, batch_size=batch_size, num_workers=20, shuffle=False, drop_last=False, pin_memory=True\n", + ")" ] }, { @@ -777,9 +791,7 @@ "plt.figure(figsize=(16, 8))\n", "plt.axis(\"off\")\n", "plt.title(\"Training Images\")\n", - "plt.imshow( \n", - " vutils.make_grid(batch[0][:16], padding=2, normalize=True).cpu().numpy().transpose((1, 2, 0))\n", - ")\n", + "plt.imshow(vutils.make_grid(batch[0][:16], padding=2, normalize=True).cpu().numpy().transpose((1, 2, 0)))\n", "\n", "batch = None\n", "torch.cuda.empty_cache()" @@ -885,20 +897,22 @@ "\n", "lr = 0.01\n", "\n", - "optimizer = optim.SGD([\n", - " {\n", - " \"params\": chain(model.stem.parameters(), model.blocks.parameters()),\n", - " \"lr\": lr * 0.1,\n", - " },\n", - " {\n", - " \"params\": model.head[:6].parameters(),\n", - " \"lr\": lr * 0.2,\n", - " }, \n", - " {\n", - " \"params\": model.head[6].parameters(), \n", - " \"lr\": lr\n", - " }], \n", - " momentum=0.9, weight_decay=0.001, nesterov=True)\n" + "optimizer = optim.SGD(\n", + " [\n", + " {\n", + " \"params\": chain(model.stem.parameters(), model.blocks.parameters()),\n", + " \"lr\": lr * 0.1,\n", + " },\n", + " {\n", + " \"params\": model.head[:6].parameters(),\n", + " \"lr\": lr * 0.2,\n", + " },\n", + " {\"params\": model.head[6].parameters(), \"lr\": lr},\n", + " ],\n", + " momentum=0.9,\n", + " weight_decay=0.001,\n", + " nesterov=True,\n", + ")" ] }, { @@ -925,7 +939,7 @@ "\n", "\n", "# Initialize Amp\n", - "model, optimizer = amp.initialize(model, optimizer, opt_level=\"O2\", num_losses=1)\n" + "model, optimizer = amp.initialize(model, optimizer, opt_level=\"O2\", num_losses=1)" ] }, { @@ -949,20 +963,20 @@ "\n", " x = convert_tensor(batch[0], device=device, non_blocking=True)\n", " y = convert_tensor(batch[1], device=device, non_blocking=True)\n", - " \n", + "\n", " y_pred = model(x)\n", - " \n", - " # Compute loss \n", - " loss = criterion(y_pred, y) \n", + "\n", + " # Compute loss\n", + " loss = criterion(y_pred, y)\n", "\n", " with amp.scale_loss(loss, optimizer) as scaled_loss:\n", " scaled_loss.backward()\n", "\n", " optimizer.step()\n", - " \n", + "\n", " return {\n", " \"batchloss\": loss.item(),\n", - " } " + " }" ] }, { @@ -1021,7 +1035,7 @@ "\n", "\n", "def output_transform(out):\n", - " return out['batchloss']\n", + " return out[\"batchloss\"]\n", "\n", "\n", "RunningAverage(output_transform=output_transform).attach(trainer, \"batchloss\")" @@ -1040,9 +1054,16 @@ "tb_logger = TensorboardLogger(log_dir=log_path)\n", "\n", "\n", - "tb_logger.attach(trainer, \n", - " log_handler=OutputHandler('training', ['batchloss', ]), \n", - " event_name=Events.ITERATION_COMPLETED)\n", + "tb_logger.attach(\n", + " trainer,\n", + " log_handler=OutputHandler(\n", + " \"training\",\n", + " [\n", + " \"batchloss\",\n", + " ],\n", + " ),\n", + " event_name=Events.ITERATION_COMPLETED,\n", + ")\n", "\n", "print(\"Experiment name: \", exp_name)" ] @@ -1063,9 +1084,7 @@ "trainer.add_event_handler(Events.EPOCH_COMPLETED, lambda engine: lr_scheduler.step())\n", "\n", "# Log optimizer parameters\n", - "tb_logger.attach(trainer,\n", - " log_handler=OptimizerParamsHandler(optimizer, \"lr\"), \n", - " event_name=Events.EPOCH_STARTED)" + "tb_logger.attach(trainer, log_handler=OptimizerParamsHandler(optimizer, \"lr\"), event_name=Events.EPOCH_STARTED)" ] }, { @@ -1080,9 +1099,9 @@ "# ProgressBar(bar_format=\"\").attach(trainer, metric_names=['batchloss',])\n", "\n", "# Epoch-wise progress bar with display of training losses\n", - "ProgressBar(persist=True, bar_format=\"\").attach(trainer, \n", - " event_name=Events.EPOCH_STARTED, \n", - " closing_event_name=Events.COMPLETED)" + "ProgressBar(persist=True, bar_format=\"\").attach(\n", + " trainer, event_name=Events.EPOCH_STARTED, closing_event_name=Events.COMPLETED\n", + ")" ] }, { @@ -1099,19 +1118,17 @@ "outputs": [], "source": [ "metrics = {\n", - " 'Loss': Loss(criterion),\n", - " 'Accuracy': Accuracy(),\n", - " 'Precision': Precision(average=True),\n", - " 'Recall': Recall(average=True),\n", - " 'Top-5 Accuracy': TopKCategoricalAccuracy(k=5)\n", + " \"Loss\": Loss(criterion),\n", + " \"Accuracy\": Accuracy(),\n", + " \"Precision\": Precision(average=True),\n", + " \"Recall\": Recall(average=True),\n", + " \"Top-5 Accuracy\": TopKCategoricalAccuracy(k=5),\n", "}\n", "\n", "\n", - "evaluator = create_supervised_evaluator(model, metrics=metrics, \n", - " device=device, non_blocking=True)\n", + "evaluator = create_supervised_evaluator(model, metrics=metrics, device=device, non_blocking=True)\n", "\n", - "train_evaluator = create_supervised_evaluator(model, metrics=metrics, \n", - " device=device, non_blocking=True)" + "train_evaluator = create_supervised_evaluator(model, metrics=metrics, device=device, non_blocking=True)" ] }, { @@ -1138,7 +1155,7 @@ " event_name=Events.EPOCH_COMPLETED,\n", " tag=\"training\",\n", " metric_names=list(metrics.keys()),\n", - " global_step_transform=global_step_from_engine(trainer)\n", + " global_step_transform=global_step_from_engine(trainer),\n", ")\n", "\n", "# Log val metrics:\n", @@ -1147,7 +1164,7 @@ " event_name=Events.EPOCH_COMPLETED,\n", " tag=\"test\",\n", " metric_names=list(metrics.keys()),\n", - " global_step_transform=global_step_from_engine(trainer)\n", + " global_step_transform=global_step_from_engine(trainer),\n", ")" ] }, @@ -1166,6 +1183,7 @@ "source": [ "import logging\n", "\n", + "\n", "# Setup engine & logger\n", "def setup_logger(logger):\n", " handler = logging.StreamHandler()\n", @@ -1189,15 +1207,15 @@ "\n", "# Store the best model\n", "def default_score_fn(engine):\n", - " score = engine.state.metrics['Accuracy']\n", + " score = engine.state.metrics[\"Accuracy\"]\n", " return score\n", "\n", + "\n", "# Force filename to model.pt to ease the rerun of the notebook\n", "disk_saver = DiskSaver(dirname=log_path)\n", - "best_model_handler = Checkpoint(to_save={'model': model}, \n", - " save_handler=disk_saver, \n", - " filename_pattern=\"{name}.{ext}\", \n", - " n_saved=1)\n", + "best_model_handler = Checkpoint(\n", + " to_save={\"model\": model}, save_handler=disk_saver, filename_pattern=\"{name}.{ext}\", n_saved=1\n", + ")\n", "evaluator.add_event_handler(Events.COMPLETED, best_model_handler)\n", "\n", "# Add early stopping\n", @@ -1211,6 +1229,7 @@ "def empty_cuda_cache(engine):\n", " torch.cuda.empty_cache()\n", " import gc\n", + "\n", " gc.collect()\n", "\n", "\n", @@ -1298,19 +1317,19 @@ "outputs": [], "source": [ "metrics = {\n", - " 'Accuracy': Accuracy(),\n", - " 'Precision': Precision(average=True),\n", - " 'Recall': Recall(average=True),\n", + " \"Accuracy\": Accuracy(),\n", + " \"Precision\": Precision(average=True),\n", + " \"Recall\": Recall(average=True),\n", "}\n", "\n", "\n", "def inference_update_with_tta(engine, batch):\n", " best_model.eval()\n", " with torch.no_grad():\n", - " x, y = batch \n", + " x, y = batch\n", " # Let's compute final prediction as a mean of predictions on x and flipped x\n", " y_pred1 = best_model(x)\n", - " y_pred2 = best_model(x.flip(dims=(-1, )))\n", + " y_pred2 = best_model(x.flip(dims=(-1,)))\n", " y_pred = 0.5 * (y_pred1 + y_pred2)\n", "\n", " return y_pred, y\n", diff --git a/examples/notebooks/FashionMNIST.ipynb b/examples/notebooks/FashionMNIST.ipynb index 0d298c1c2d0c..624e53b84204 100644 --- a/examples/notebooks/FashionMNIST.ipynb +++ b/examples/notebooks/FashionMNIST.ipynb @@ -126,15 +126,14 @@ "outputs": [], "source": [ "# transform to normalize the data\n", - "transform = transforms.Compose([transforms.ToTensor(),\n", - " transforms.Normalize((0.5,), (0.5,))])\n", + "transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,))])\n", "\n", "# Download and load the training data\n", - "trainset = datasets.FashionMNIST('./data', download=True, train=True, transform=transform)\n", + "trainset = datasets.FashionMNIST(\"./data\", download=True, train=True, transform=transform)\n", "train_loader = DataLoader(trainset, batch_size=64, shuffle=True)\n", "\n", "# Download and load the test data\n", - "validationset = datasets.FashionMNIST('./data', download=True, train=False, transform=transform)\n", + "validationset = datasets.FashionMNIST(\"./data\", download=True, train=False, transform=transform)\n", "val_loader = DataLoader(validationset, batch_size=64, shuffle=True)" ] }, @@ -164,39 +163,30 @@ "outputs": [], "source": [ "class CNN(nn.Module):\n", - " \n", " def __init__(self):\n", " super(CNN, self).__init__()\n", - " \n", + "\n", " self.convlayer1 = nn.Sequential(\n", - " nn.Conv2d(1, 32, 3,padding=1),\n", - " nn.BatchNorm2d(32),\n", - " nn.ReLU(),\n", - " nn.MaxPool2d(kernel_size=2, stride=2)\n", - " )\n", - " \n", - " self.convlayer2 = nn.Sequential(\n", - " nn.Conv2d(32,64,3),\n", - " nn.BatchNorm2d(64),\n", - " nn.ReLU(),\n", - " nn.MaxPool2d(2)\n", + " nn.Conv2d(1, 32, 3, padding=1), nn.BatchNorm2d(32), nn.ReLU(), nn.MaxPool2d(kernel_size=2, stride=2)\n", " )\n", - " \n", - " self.fc1 = nn.Linear(64*6*6,600)\n", + "\n", + " self.convlayer2 = nn.Sequential(nn.Conv2d(32, 64, 3), nn.BatchNorm2d(64), nn.ReLU(), nn.MaxPool2d(2))\n", + "\n", + " self.fc1 = nn.Linear(64 * 6 * 6, 600)\n", " self.drop = nn.Dropout2d(0.25)\n", " self.fc2 = nn.Linear(600, 120)\n", " self.fc3 = nn.Linear(120, 10)\n", - " \n", + "\n", " def forward(self, x):\n", " x = self.convlayer1(x)\n", " x = self.convlayer2(x)\n", - " x = x.view(-1,64*6*6)\n", + " x = x.view(-1, 64 * 6 * 6)\n", " x = self.fc1(x)\n", " x = self.drop(x)\n", " x = self.fc2(x)\n", " x = self.fc3(x)\n", - " \n", - " return F.log_softmax(x,dim=1)" + "\n", + " return F.log_softmax(x, dim=1)" ] }, { @@ -260,15 +250,11 @@ "epochs = 12\n", "# creating trainer,evaluator\n", "trainer = create_supervised_trainer(model, optimizer, criterion, device=device)\n", - "metrics = {\n", - " 'accuracy':Accuracy(),\n", - " 'nll':Loss(criterion),\n", - " 'cm':ConfusionMatrix(num_classes=10)\n", - "}\n", + "metrics = {\"accuracy\": Accuracy(), \"nll\": Loss(criterion), \"cm\": ConfusionMatrix(num_classes=10)}\n", "train_evaluator = create_supervised_evaluator(model, metrics=metrics, device=device)\n", "val_evaluator = create_supervised_evaluator(model, metrics=metrics, device=device)\n", - "training_history = {'accuracy':[],'loss':[]}\n", - "validation_history = {'accuracy':[],'loss':[]}\n", + "training_history = {\"accuracy\": [], \"loss\": []}\n", + "validation_history = {\"accuracy\": [], \"loss\": []}\n", "last_epoch = []" ] }, @@ -287,7 +273,7 @@ "metadata": {}, "outputs": [], "source": [ - "RunningAverage(output_transform=lambda x: x).attach(trainer, 'loss')" + "RunningAverage(output_transform=lambda x: x).attach(trainer, \"loss\")" ] }, { @@ -306,9 +292,10 @@ "outputs": [], "source": [ "def score_function(engine):\n", - " val_loss = engine.state.metrics['nll']\n", + " val_loss = engine.state.metrics[\"nll\"]\n", " return -val_loss\n", "\n", + "\n", "handler = EarlyStopping(patience=10, score_function=score_function, trainer=trainer)\n", "val_evaluator.add_event_handler(Events.COMPLETED, handler)" ] @@ -338,25 +325,33 @@ "def log_training_results(trainer):\n", " train_evaluator.run(train_loader)\n", " metrics = train_evaluator.state.metrics\n", - " accuracy = metrics['accuracy']*100\n", - " loss = metrics['nll']\n", + " accuracy = metrics[\"accuracy\"] * 100\n", + " loss = metrics[\"nll\"]\n", " last_epoch.append(0)\n", - " training_history['accuracy'].append(accuracy)\n", - " training_history['loss'].append(loss)\n", - " print(\"Training Results - Epoch: {} Avg accuracy: {:.2f} Avg loss: {:.2f}\"\n", - " .format(trainer.state.epoch, accuracy, loss))\n", + " training_history[\"accuracy\"].append(accuracy)\n", + " training_history[\"loss\"].append(loss)\n", + " print(\n", + " \"Training Results - Epoch: {} Avg accuracy: {:.2f} Avg loss: {:.2f}\".format(\n", + " trainer.state.epoch, accuracy, loss\n", + " )\n", + " )\n", + "\n", "\n", "def log_validation_results(trainer):\n", " val_evaluator.run(val_loader)\n", " metrics = val_evaluator.state.metrics\n", - " accuracy = metrics['accuracy']*100\n", - " loss = metrics['nll']\n", - " validation_history['accuracy'].append(accuracy)\n", - " validation_history['loss'].append(loss)\n", - " print(\"Validation Results - Epoch: {} Avg accuracy: {:.2f} Avg loss: {:.2f}\"\n", - " .format(trainer.state.epoch, accuracy, loss))\n", - " \n", - "trainer.add_event_handler(Events.EPOCH_COMPLETED, log_validation_results) " + " accuracy = metrics[\"accuracy\"] * 100\n", + " loss = metrics[\"nll\"]\n", + " validation_history[\"accuracy\"].append(accuracy)\n", + " validation_history[\"loss\"].append(loss)\n", + " print(\n", + " \"Validation Results - Epoch: {} Avg accuracy: {:.2f} Avg loss: {:.2f}\".format(\n", + " trainer.state.epoch, accuracy, loss\n", + " )\n", + " )\n", + "\n", + "\n", + "trainer.add_event_handler(Events.EPOCH_COMPLETED, log_validation_results)" ] }, { @@ -385,19 +380,19 @@ "def log_confusion_matrix(trainer):\n", " val_evaluator.run(val_loader)\n", " metrics = val_evaluator.state.metrics\n", - " cm = metrics['cm']\n", + " cm = metrics[\"cm\"]\n", " cm = cm.numpy()\n", " cm = cm.astype(int)\n", - " classes = ['T-shirt/top','Trouser','Pullover','Dress','Coat','Sandal','Shirt','Sneaker','Bag','Ankle Boot']\n", - " fig, ax = plt.subplots(figsize=(10,10)) \n", - " ax= plt.subplot()\n", - " sns.heatmap(cm, annot=True, ax = ax,fmt=\"d\")\n", + " classes = [\"T-shirt/top\", \"Trouser\", \"Pullover\", \"Dress\", \"Coat\", \"Sandal\", \"Shirt\", \"Sneaker\", \"Bag\", \"Ankle Boot\"]\n", + " fig, ax = plt.subplots(figsize=(10, 10))\n", + " ax = plt.subplot()\n", + " sns.heatmap(cm, annot=True, ax=ax, fmt=\"d\")\n", " # labels, title and ticks\n", - " ax.set_xlabel('Predicted labels')\n", - " ax.set_ylabel('True labels') \n", - " ax.set_title('Confusion Matrix') \n", - " ax.xaxis.set_ticklabels(classes,rotation=90)\n", - " ax.yaxis.set_ticklabels(classes,rotation=0)" + " ax.set_xlabel(\"Predicted labels\")\n", + " ax.set_ylabel(\"True labels\")\n", + " ax.set_title(\"Confusion Matrix\")\n", + " ax.xaxis.set_ticklabels(classes, rotation=90)\n", + " ax.yaxis.set_ticklabels(classes, rotation=0)" ] }, { @@ -417,8 +412,8 @@ "metadata": {}, "outputs": [], "source": [ - "checkpointer = ModelCheckpoint('./saved_models', 'fashionMNIST', n_saved=2, create_dir=True, require_empty=False)\n", - "trainer.add_event_handler(Events.EPOCH_COMPLETED, checkpointer, {'fashionMNIST': model})" + "checkpointer = ModelCheckpoint(\"./saved_models\", \"fashionMNIST\", n_saved=2, create_dir=True, require_empty=False)\n", + "trainer.add_event_handler(Events.EPOCH_COMPLETED, checkpointer, {\"fashionMNIST\": model})" ] }, { @@ -453,10 +448,10 @@ "metadata": {}, "outputs": [], "source": [ - "plt.plot(training_history['accuracy'],label=\"Training Accuracy\")\n", - "plt.plot(validation_history['accuracy'],label=\"Validation Accuracy\")\n", - "plt.xlabel('No. of Epochs')\n", - "plt.ylabel('Accuracy')\n", + "plt.plot(training_history[\"accuracy\"], label=\"Training Accuracy\")\n", + "plt.plot(validation_history[\"accuracy\"], label=\"Validation Accuracy\")\n", + "plt.xlabel(\"No. of Epochs\")\n", + "plt.ylabel(\"Accuracy\")\n", "plt.legend(frameon=False)\n", "plt.show()" ] @@ -467,10 +462,10 @@ "metadata": {}, "outputs": [], "source": [ - "plt.plot(training_history['loss'],label=\"Training Loss\")\n", - "plt.plot(validation_history['loss'],label=\"Validation Loss\")\n", - "plt.xlabel('No. of Epochs')\n", - "plt.ylabel('Loss')\n", + "plt.plot(training_history[\"loss\"], label=\"Training Loss\")\n", + "plt.plot(validation_history[\"loss\"], label=\"Validation Loss\")\n", + "plt.xlabel(\"No. of Epochs\")\n", + "plt.ylabel(\"Loss\")\n", "plt.legend(frameon=False)\n", "plt.show()" ] @@ -493,15 +488,15 @@ "def fetch_last_checkpoint_model_filename(model_save_path):\n", " import os\n", " from pathlib import Path\n", + "\n", " checkpoint_files = os.listdir(model_save_path)\n", - " checkpoint_files = [f for f in checkpoint_files if '.pt' in f]\n", - " checkpoint_iter = [\n", - " int(x.split('_')[2].split('.')[0])\n", - " for x in checkpoint_files]\n", + " checkpoint_files = [f for f in checkpoint_files if \".pt\" in f]\n", + " checkpoint_iter = [int(x.split(\"_\")[2].split(\".\")[0]) for x in checkpoint_files]\n", " last_idx = np.array(checkpoint_iter).argmax()\n", " return Path(model_save_path) / checkpoint_files[last_idx]\n", "\n", - "model.load_state_dict(torch.load(fetch_last_checkpoint_model_filename('./saved_models')))\n", + "\n", + "model.load_state_dict(torch.load(fetch_last_checkpoint_model_filename(\"./saved_models\")))\n", "print(\"Model Loaded\")" ] }, @@ -522,29 +517,31 @@ "outputs": [], "source": [ "# classes of fashion mnist dataset\n", - "classes = ['T-shirt/top','Trouser','Pullover','Dress','Coat','Sandal','Shirt','Sneaker','Bag','Ankle Boot']\n", + "classes = [\"T-shirt/top\", \"Trouser\", \"Pullover\", \"Dress\", \"Coat\", \"Sandal\", \"Shirt\", \"Sneaker\", \"Bag\", \"Ankle Boot\"]\n", "# creating iterator for iterating the dataset\n", "dataiter = iter(val_loader)\n", "images, labels = next(dataiter)\n", "images_arr = []\n", "labels_arr = []\n", "pred_arr = []\n", - "# moving model to cpu for inference \n", + "# moving model to cpu for inference\n", "model.to(\"cpu\")\n", "# iterating on the dataset to predict the output\n", - "for i in range(0,10):\n", + "for i in range(0, 10):\n", " images_arr.append(images[i].unsqueeze(0))\n", " labels_arr.append(labels[i].item())\n", " ps = torch.exp(model(images_arr[i]))\n", " ps = ps.data.numpy().squeeze()\n", " pred_arr.append(np.argmax(ps))\n", "# plotting the results\n", - "fig = plt.figure(figsize=(25,4))\n", + "fig = plt.figure(figsize=(25, 4))\n", "for i in range(10):\n", - " ax = fig.add_subplot(2, 20//2, i+1, xticks=[], yticks=[])\n", + " ax = fig.add_subplot(2, 20 // 2, i + 1, xticks=[], yticks=[])\n", " ax.imshow(images_arr[i].resize_(1, 28, 28).numpy().squeeze())\n", - " ax.set_title(\"{} ({})\".format(classes[pred_arr[i]], classes[labels_arr[i]]),\n", - " color=(\"green\" if pred_arr[i]==labels_arr[i] else \"red\"))" + " ax.set_title(\n", + " \"{} ({})\".format(classes[pred_arr[i]], classes[labels_arr[i]]),\n", + " color=(\"green\" if pred_arr[i] == labels_arr[i] else \"red\"),\n", + " )" ] }, { diff --git a/examples/notebooks/FastaiLRFinder_MNIST.ipynb b/examples/notebooks/FastaiLRFinder_MNIST.ipynb index 27e7052ff30e..55f09f15d084 100644 --- a/examples/notebooks/FastaiLRFinder_MNIST.ipynb +++ b/examples/notebooks/FastaiLRFinder_MNIST.ipynb @@ -53,6 +53,7 @@ "outputs": [], "source": [ "import ignite\n", + "\n", "ignite.__version__" ] }, @@ -82,7 +83,7 @@ "outputs": [], "source": [ "mnist_pwd = \"data\"\n", - "batch_size= 256" + "batch_size = 256" ] }, { @@ -169,7 +170,7 @@ "ProgressBar(persist=True).attach(trainer, output_transform=lambda x: {\"batch loss\": x})\n", "\n", "lr_finder = FastaiLRFinder()\n", - "to_save={'model': model, 'optimizer': optimizer}\n", + "to_save = {\"model\": model, \"optimizer\": optimizer}\n", "with lr_finder.attach(trainer, to_save, diverge_th=1.5) as trainer_with_lr_finder:\n", " trainer_with_lr_finder.run(trainloader)" ] @@ -188,6 +189,7 @@ "outputs": [], "source": [ "from matplotlib import pyplot as plt\n", + "\n", "ax = lr_finder.plot()\n", "plt.show()\n", "\n", @@ -257,7 +259,7 @@ "outputs": [], "source": [ "lr_finder.apply_suggested_lr(optimizer)\n", - "print(optimizer.param_groups[0]['lr'])" + "print(optimizer.param_groups[0][\"lr\"])" ] }, { diff --git a/examples/notebooks/HandlersTimeProfiler_MNIST.ipynb b/examples/notebooks/HandlersTimeProfiler_MNIST.ipynb index f65147479b43..9263590b9306 100644 --- a/examples/notebooks/HandlersTimeProfiler_MNIST.ipynb +++ b/examples/notebooks/HandlersTimeProfiler_MNIST.ipynb @@ -53,6 +53,7 @@ "# A hack to fix the horizontal spill in large output\n", "# ref: https://stackoverflow.com/a/59058418/6574605\n", "from IPython.core.display import HTML\n", + "\n", "display(HTML(\"\"))" ] }, @@ -81,7 +82,7 @@ "outputs": [], "source": [ "mnist_pwd = \"data\"\n", - "batch_size= 256" + "batch_size = 256" ] }, { @@ -168,6 +169,7 @@ "pbar = ProgressBar(persist=True)\n", "pbar.attach(trainer, metric_names=\"all\")\n", "\n", + "\n", "# Evaluate on each epoch using event handler\n", "@trainer.on(Events.EPOCH_COMPLETED)\n", "def log_validation_results(engine):\n", diff --git a/examples/notebooks/TextTransformers.ipynb b/examples/notebooks/TextTransformers.ipynb index 04f583627696..cd4b386d62d4 100644 --- a/examples/notebooks/TextTransformers.ipynb +++ b/examples/notebooks/TextTransformers.ipynb @@ -92,10 +92,12 @@ "checkpoint = \"distilbert-base-uncased\"\n", "tokenizer = AutoTokenizer.from_pretrained(checkpoint)\n", "\n", + "\n", "# 3. Tokenize function\n", "def tokenize_function(example):\n", " return tokenizer(example[\"text\"], truncation=True)\n", "\n", + "\n", "# 4. Apply tokenization to the entire dataset\n", "tokenized_datasets = raw_datasets.map(tokenize_function, batched=True)\n", "\n", @@ -109,12 +111,8 @@ "data_collator = DataCollatorWithPadding(tokenizer=tokenizer)\n", "\n", "# 7. Create DataLoaders\n", - "train_dataloader = DataLoader(\n", - " tokenized_datasets[\"train\"], shuffle=True, batch_size=32, collate_fn=data_collator\n", - ")\n", - "test_dataloader = DataLoader(\n", - " tokenized_datasets[\"test\"], batch_size=32, collate_fn=data_collator\n", - ")" + "train_dataloader = DataLoader(tokenized_datasets[\"train\"], shuffle=True, batch_size=32, collate_fn=data_collator)\n", + "test_dataloader = DataLoader(tokenized_datasets[\"test\"], batch_size=32, collate_fn=data_collator)" ] }, { @@ -165,6 +163,7 @@ "# Using Mixed Precision for speed\n", "scaler = GradScaler()\n", "\n", + "\n", "def process_function(engine, batch):\n", " model.train()\n", " optimizer.zero_grad()\n", @@ -173,7 +172,7 @@ " batch = {k: v.to(device) for k, v in batch.items()}\n", "\n", " # Forward pass with AMP\n", - " with torch.amp.autocast('cuda'):\n", + " with torch.amp.autocast(\"cuda\"):\n", " outputs = model(**batch)\n", " loss = outputs.loss\n", "\n", @@ -184,6 +183,7 @@ "\n", " return loss.item()\n", "\n", + "\n", "def eval_function(engine, batch):\n", " model.eval()\n", " with torch.no_grad():\n", @@ -194,6 +194,7 @@ " # Return (y_pred, y) for Ignite metrics to consume\n", " return logits, batch[\"labels\"]\n", "\n", + "\n", "# Instantiate Engines\n", "trainer = Engine(process_function)\n", "train_evaluator = Engine(eval_function)\n", @@ -224,13 +225,10 @@ "outputs": [], "source": [ "# 1. Running Average of Loss\n", - "RunningAverage(output_transform=lambda x: x).attach(trainer, 'loss')\n", + "RunningAverage(output_transform=lambda x: x).attach(trainer, \"loss\")\n", "\n", "# 2. Accuracy and Loss Metrics\n", - "metrics = {\n", - " 'accuracy': Accuracy(),\n", - " 'nll': Loss(torch.nn.CrossEntropyLoss())\n", - "}\n", + "metrics = {\"accuracy\": Accuracy(), \"nll\": Loss(torch.nn.CrossEntropyLoss())}\n", "\n", "for name, metric in metrics.items():\n", " metric.attach(train_evaluator, name)\n", @@ -238,43 +236,45 @@ "\n", "# 3. Progress Bar\n", "pbar = ProgressBar(persist=True, bar_format=\"\")\n", - "pbar.attach(trainer, ['loss'])\n", + "pbar.attach(trainer, [\"loss\"])\n", "\n", "eval_pbar = ProgressBar(desc=\"Evaluating\", persist=False)\n", "eval_pbar.attach(validation_evaluator)\n", "\n", + "\n", "# 4. Log Validation Results at the end of every epoch\n", "@trainer.on(Events.EPOCH_COMPLETED)\n", "def log_validation_results(engine):\n", " validation_evaluator.run(test_dataloader)\n", " metrics = validation_evaluator.state.metrics\n", - " avg_accuracy = metrics['accuracy']\n", - " avg_nll = metrics['nll']\n", + " avg_accuracy = metrics[\"accuracy\"]\n", + " avg_nll = metrics[\"nll\"]\n", "\n", " pbar.log_message(\n", - " f\"Validation Results - Epoch: {engine.state.epoch} \"\n", - " f\"Avg accuracy: {avg_accuracy:.2f} Avg loss: {avg_nll:.2f}\"\n", + " f\"Validation Results - Epoch: {engine.state.epoch} Avg accuracy: {avg_accuracy:.2f} Avg loss: {avg_nll:.2f}\"\n", " )\n", "\n", + "\n", "# 5. Early Stopping\n", "def score_function(engine):\n", - " return engine.state.metrics['accuracy']\n", + " return engine.state.metrics[\"accuracy\"]\n", + "\n", "\n", "handler = EarlyStopping(patience=2, score_function=score_function, trainer=trainer)\n", "validation_evaluator.add_event_handler(Events.COMPLETED, handler)\n", "\n", "# 6. Model Checkpoint\n", - "to_save = {'model': model}\n", - "save_handler = DiskSaver(dirname='/tmp/models', create_dir=True, require_empty=False)\n", + "to_save = {\"model\": model}\n", + "save_handler = DiskSaver(dirname=\"/tmp/models\", create_dir=True, require_empty=False)\n", "\n", "checkpointer = Checkpoint(\n", " to_save=to_save,\n", " save_handler=save_handler,\n", - " filename_prefix='distilbert_imdb',\n", + " filename_prefix=\"distilbert_imdb\",\n", " n_saved=1,\n", " score_function=score_function,\n", " score_name=\"val_acc\",\n", - " global_step_transform=global_step_from_engine(trainer)\n", + " global_step_transform=global_step_from_engine(trainer),\n", ")\n", "\n", "validation_evaluator.add_event_handler(Events.COMPLETED, checkpointer)" @@ -341,8 +341,7 @@ "predicted_label = label_mapping[predicted_class_id]\n", "\n", "print(f\"Review: '{test_review}'\")\n", - "print(f\"Predicted Sentiment: {predicted_label}\")\n", - "\n" + "print(f\"Predicted Sentiment: {predicted_label}\")" ] } ], diff --git a/examples/notebooks/VAE.ipynb b/examples/notebooks/VAE.ipynb index fa28da204e02..1467968c9c7e 100644 --- a/examples/notebooks/VAE.ipynb +++ b/examples/notebooks/VAE.ipynb @@ -89,6 +89,7 @@ "from torch.utils.data import DataLoader\n", "from torch import nn, optim\n", "from torch.nn import functional as F\n", + "\n", "SEED = 1234\n", "\n", "torch.manual_seed(SEED)\n", @@ -176,12 +177,12 @@ "image = train_data[0][0]\n", "label = train_data[0][1]\n", "\n", - "print ('len(train_data) : ', len(train_data))\n", - "print ('len(val_data) : ', len(val_data))\n", - "print ('image.shape : ', image.shape)\n", - "print ('label : ', label)\n", + "print(\"len(train_data) : \", len(train_data))\n", + "print(\"len(val_data) : \", len(val_data))\n", + "print(\"image.shape : \", image.shape)\n", + "print(\"label : \", label)\n", "\n", - "img = plt.imshow(image.squeeze().numpy(), cmap='gray')" + "img = plt.imshow(image.squeeze().numpy(), cmap=\"gray\")" ] }, { @@ -202,7 +203,7 @@ "metadata": {}, "outputs": [], "source": [ - "kwargs = {'num_workers': 1, 'pin_memory': True} if device == 'cuda' else {}\n", + "kwargs = {\"num_workers\": 1, \"pin_memory\": True} if device == \"cuda\" else {}\n", "\n", "train_loader = DataLoader(train_data, batch_size=32, shuffle=True, **kwargs)\n", "val_loader = DataLoader(val_data, batch_size=32, shuffle=True, **kwargs)\n", @@ -211,8 +212,8 @@ " x, y = batch\n", " break\n", "\n", - "print ('x.shape : ', x.shape)\n", - "print ('y.shape : ', y.shape)" + "print(\"x.shape : \", x.shape)\n", + "print(\"y.shape : \", y.shape)" ] }, { @@ -265,7 +266,7 @@ " return self.fc21(h1), self.fc22(h1)\n", "\n", " def reparameterize(self, mu, logvar):\n", - " std = torch.exp(0.5*logvar)\n", + " std = torch.exp(0.5 * logvar)\n", " eps = torch.randn_like(std)\n", " return eps.mul(std).add_(mu)\n", "\n", @@ -302,6 +303,7 @@ "model = VAE().to(device)\n", "optimizer = optim.Adam(model.parameters(), lr=1e-3)\n", "\n", + "\n", "def kld_loss(x_pred, x, mu, logvar):\n", " # see Appendix B from VAE paper:\n", " # Kingma and Welling. Auto-Encoding Variational Bayes. ICLR, 2014\n", @@ -309,7 +311,8 @@ " # 0.5 * sum(1 + log(sigma^2) - mu^2 - sigma^2)\n", " return -0.5 * torch.sum(1 + logvar - mu.pow(2) - logvar.exp())\n", "\n", - "bce_loss = nn.BCELoss(reduction='sum')" + "\n", + "bce_loss = nn.BCELoss(reduction=\"sum\")" ] }, { @@ -397,7 +400,7 @@ " x = x.to(device)\n", " x = x.view(-1, 784)\n", " x_pred, mu, logvar = model(x)\n", - " kwargs = {'mu': mu, 'logvar': logvar}\n", + " kwargs = {\"mu\": mu, \"logvar\": logvar}\n", " return x_pred, x, kwargs" ] }, @@ -418,8 +421,8 @@ "source": [ "trainer = Engine(process_function)\n", "evaluator = Engine(evaluate_function)\n", - "training_history = {'bce': [], 'kld': [], 'mse': []}\n", - "validation_history = {'bce': [], 'kld': [], 'mse': []}" + "training_history = {\"bce\": [], \"kld\": [], \"mse\": []}\n", + "validation_history = {\"bce\": [], \"kld\": [], \"mse\": []}" ] }, { @@ -437,9 +440,9 @@ "metadata": {}, "outputs": [], "source": [ - "RunningAverage(output_transform=lambda x: x[0]).attach(trainer, 'loss')\n", - "RunningAverage(output_transform=lambda x: x[1]).attach(trainer, 'bce')\n", - "RunningAverage(output_transform=lambda x: x[2]).attach(trainer, 'kld')" + "RunningAverage(output_transform=lambda x: x[0]).attach(trainer, \"loss\")\n", + "RunningAverage(output_transform=lambda x: x[1]).attach(trainer, \"bce\")\n", + "RunningAverage(output_transform=lambda x: x[2]).attach(trainer, \"kld\")" ] }, { @@ -459,9 +462,9 @@ "metadata": {}, "outputs": [], "source": [ - "MeanSquaredError(output_transform=lambda x: [x[0], x[1]]).attach(evaluator, 'mse')\n", - "Loss(bce_loss, output_transform=lambda x: [x[0], x[1]]).attach(evaluator, 'bce')\n", - "Loss(kld_loss).attach(evaluator, 'kld')" + "MeanSquaredError(output_transform=lambda x: [x[0], x[1]]).attach(evaluator, \"mse\")\n", + "Loss(bce_loss, output_transform=lambda x: [x[0], x[1]]).attach(evaluator, \"bce\")\n", + "Loss(kld_loss).attach(evaluator, \"kld\")" ] }, { @@ -488,11 +491,14 @@ "source": [ "@trainer.on(Events.EPOCH_COMPLETED)\n", "def print_trainer_logs(engine):\n", - " avg_loss = engine.state.metrics['loss']\n", - " avg_bce = engine.state.metrics['bce']\n", - " avg_kld = engine.state.metrics['kld']\n", - " print(\"Trainer Results - Epoch {} - Avg loss: {:.2f} Avg bce: {:.2f} Avg kld: {:.2f}\"\n", - " .format(engine.state.epoch, avg_loss, avg_bce, avg_kld))" + " avg_loss = engine.state.metrics[\"loss\"]\n", + " avg_bce = engine.state.metrics[\"bce\"]\n", + " avg_kld = engine.state.metrics[\"kld\"]\n", + " print(\n", + " \"Trainer Results - Epoch {} - Avg loss: {:.2f} Avg bce: {:.2f} Avg kld: {:.2f}\".format(\n", + " engine.state.epoch, avg_loss, avg_bce, avg_kld\n", + " )\n", + " )" ] }, { @@ -511,18 +517,22 @@ "def print_logs(engine, dataloader, mode, history_dict):\n", " evaluator.run(dataloader, max_epochs=1)\n", " metrics = evaluator.state.metrics\n", - " avg_mse = metrics['mse']\n", - " avg_bce = metrics['bce']\n", - " avg_kld = metrics['kld']\n", - " avg_loss = avg_bce + avg_kld\n", + " avg_mse = metrics[\"mse\"]\n", + " avg_bce = metrics[\"bce\"]\n", + " avg_kld = metrics[\"kld\"]\n", + " avg_loss = avg_bce + avg_kld\n", " print(\n", - " mode + \" Results - Epoch {} - Avg mse: {:.2f} Avg loss: {:.2f} Avg bce: {:.2f} Avg kld: {:.2f}\"\n", - " .format(engine.state.epoch, avg_mse, avg_loss, avg_bce, avg_kld))\n", + " mode\n", + " + \" Results - Epoch {} - Avg mse: {:.2f} Avg loss: {:.2f} Avg bce: {:.2f} Avg kld: {:.2f}\".format(\n", + " engine.state.epoch, avg_mse, avg_loss, avg_bce, avg_kld\n", + " )\n", + " )\n", " for key in evaluator.state.metrics.keys():\n", " history_dict[key].append(evaluator.state.metrics[key])\n", "\n", - "trainer.add_event_handler(Events.EPOCH_COMPLETED, print_logs, train_loader, 'Training', training_history)\n", - "trainer.add_event_handler(Events.EPOCH_COMPLETED, print_logs, val_loader, 'Validation', validation_history)" + "\n", + "trainer.add_event_handler(Events.EPOCH_COMPLETED, print_logs, train_loader, \"Training\", training_history)\n", + "trainer.add_event_handler(Events.EPOCH_COMPLETED, print_logs, val_loader, \"Validation\", validation_history)" ] }, { @@ -543,12 +553,13 @@ " reconstructed_images = model(fixed_images.view(-1, 784))[0].view(-1, 1, 28, 28)\n", " comparison = torch.cat([fixed_images, reconstructed_images])\n", " if save_img:\n", - " save_image(comparison.detach().cpu(), 'reconstructed_epoch_' + str(epoch) + '.png', nrow=8)\n", + " save_image(comparison.detach().cpu(), \"reconstructed_epoch_\" + str(epoch) + \".png\", nrow=8)\n", " comparison_image = make_grid(comparison.detach().cpu(), nrow=8)\n", - " fig = plt.figure(figsize=(5, 5));\n", - " output = plt.imshow(comparison_image.permute(1, 2, 0));\n", - " plt.title('Epoch ' + str(epoch));\n", - " plt.show();\n", + " fig = plt.figure(figsize=(5, 5))\n", + " output = plt.imshow(comparison_image.permute(1, 2, 0))\n", + " plt.title(\"Epoch \" + str(epoch))\n", + " plt.show()\n", + "\n", "\n", "trainer.add_event_handler(Events.STARTED, compare_images, save_img=False)\n", "trainer.add_event_handler(Events.EPOCH_COMPLETED(every=5), compare_images, save_img=False)" @@ -587,12 +598,12 @@ "metadata": {}, "outputs": [], "source": [ - "plt.plot(range(20), training_history['bce'], 'dodgerblue', label='training')\n", - "plt.plot(range(20), validation_history['bce'], 'orange', label='validation')\n", - "plt.xlim(0, 20);\n", - "plt.xlabel('Epoch')\n", - "plt.ylabel('BCE')\n", - "plt.title('Binary Cross Entropy on Training/Validation Set')\n", + "plt.plot(range(20), training_history[\"bce\"], \"dodgerblue\", label=\"training\")\n", + "plt.plot(range(20), validation_history[\"bce\"], \"orange\", label=\"validation\")\n", + "plt.xlim(0, 20)\n", + "plt.xlabel(\"Epoch\")\n", + "plt.ylabel(\"BCE\")\n", + "plt.title(\"Binary Cross Entropy on Training/Validation Set\")\n", "plt.legend();" ] }, @@ -602,12 +613,12 @@ "metadata": {}, "outputs": [], "source": [ - "plt.plot(range(20), training_history['kld'], 'dodgerblue', label='training')\n", - "plt.plot(range(20), validation_history['kld'], 'orange', label='validation')\n", - "plt.xlim(0, 20);\n", - "plt.xlabel('Epoch')\n", - "plt.ylabel('KLD')\n", - "plt.title('KL Divergence on Training/Validation Set')\n", + "plt.plot(range(20), training_history[\"kld\"], \"dodgerblue\", label=\"training\")\n", + "plt.plot(range(20), validation_history[\"kld\"], \"orange\", label=\"validation\")\n", + "plt.xlim(0, 20)\n", + "plt.xlabel(\"Epoch\")\n", + "plt.ylabel(\"KLD\")\n", + "plt.title(\"KL Divergence on Training/Validation Set\")\n", "plt.legend();" ] }, @@ -617,12 +628,12 @@ "metadata": {}, "outputs": [], "source": [ - "plt.plot(range(20), training_history['mse'], 'dodgerblue', label='training')\n", - "plt.plot(range(20), validation_history['mse'], 'orange', label='validation')\n", - "plt.xlim(0, 20);\n", - "plt.xlabel('Epoch')\n", - "plt.ylabel('MSE')\n", - "plt.title('Mean Squared Error on Training/Validation Set')\n", + "plt.plot(range(20), training_history[\"mse\"], \"dodgerblue\", label=\"training\")\n", + "plt.plot(range(20), validation_history[\"mse\"], \"orange\", label=\"validation\")\n", + "plt.xlim(0, 20)\n", + "plt.xlabel(\"Epoch\")\n", + "plt.ylabel(\"MSE\")\n", + "plt.title(\"Mean Squared Error on Training/Validation Set\")\n", "plt.legend();" ] } diff --git a/examples/siamese_network/siamese_network.py b/examples/siamese_network/siamese_network.py index b4a9d22e85e2..b0d82a52e8c8 100644 --- a/examples/siamese_network/siamese_network.py +++ b/examples/siamese_network/siamese_network.py @@ -249,7 +249,7 @@ def test_step(engine, batch): @trainer.on(Events.EPOCH_COMPLETED(every=args.log_interval)) def test(engine): state = evaluator.run(test_loader) - print(f'Test Accuracy: {state.metrics["accuracy"]}') + print(f"Test Accuracy: {state.metrics['accuracy']}") # run the trainer trainer.run(train_loader, max_epochs=args.epochs) diff --git a/ignite/contrib/engines/common.py b/ignite/contrib/engines/common.py index 4dc774cdfdd8..958b1a64caf4 100644 --- a/ignite/contrib/engines/common.py +++ b/ignite/contrib/engines/common.py @@ -192,7 +192,9 @@ def _setup_common_training_handlers( if with_gpu_stats: GpuInfo().attach( - trainer, name="gpu", event_name=Events.ITERATION_COMPLETED(every=log_every_iters) # type: ignore[arg-type] + trainer, + name="gpu", + event_name=Events.ITERATION_COMPLETED(every=log_every_iters), # type: ignore[arg-type] ) if output_names is not None: diff --git a/ignite/contrib/handlers/base_logger.py b/ignite/contrib/handlers/base_logger.py index edf82a47f194..d8cc2657b8b1 100644 --- a/ignite/contrib/handlers/base_logger.py +++ b/ignite/contrib/handlers/base_logger.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.handlers.base_logger`` was moved to ``ignite.handlers.base_logger``. +"""``ignite.contrib.handlers.base_logger`` was moved to ``ignite.handlers.base_logger``. Note: ``ignite.contrib.handlers.base_logger`` was moved to ``ignite.handlers.base_logger``. Please refer to :mod:`~ignite.handlers.base_logger`. diff --git a/ignite/contrib/handlers/clearml_logger.py b/ignite/contrib/handlers/clearml_logger.py index 2c08251179ce..cc6d6cc18a06 100644 --- a/ignite/contrib/handlers/clearml_logger.py +++ b/ignite/contrib/handlers/clearml_logger.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.handlers.clearml_logger`` was moved to ``ignite.handlers.clearml_logger``. +"""``ignite.contrib.handlers.clearml_logger`` was moved to ``ignite.handlers.clearml_logger``. Note: ``ignite.contrib.handlers.clearml_logger`` was moved to ``ignite.handlers.clearml_logger``. Please refer to :mod:`~ignite.handlers.clearml_logger`. diff --git a/ignite/contrib/handlers/lr_finder.py b/ignite/contrib/handlers/lr_finder.py index b1995fea5f1c..847192a595bc 100644 --- a/ignite/contrib/handlers/lr_finder.py +++ b/ignite/contrib/handlers/lr_finder.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.handlers.lr_finder`` was moved to ``ignite.handlers.lr_finder``. +"""``ignite.contrib.handlers.lr_finder`` was moved to ``ignite.handlers.lr_finder``. Note: ``ignite.contrib.handlers.lr_finder`` was moved to ``ignite.handlers.lr_finder``. Please refer to :mod:`~ignite.handlers.lr_finder`. diff --git a/ignite/contrib/handlers/mlflow_logger.py b/ignite/contrib/handlers/mlflow_logger.py index eef8c925d6c4..bbd39635a0fc 100644 --- a/ignite/contrib/handlers/mlflow_logger.py +++ b/ignite/contrib/handlers/mlflow_logger.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.handlers.mlflow_logger`` was moved to ``ignite.handlers.mlflow_logger``. +"""``ignite.contrib.handlers.mlflow_logger`` was moved to ``ignite.handlers.mlflow_logger``. Note: ``ignite.contrib.handlers.mlflow_logger`` was moved to ``ignite.handlers.mlflow_logger``. Please refer to :mod:`~ignite.handlers.mlflow_logger`. diff --git a/ignite/contrib/handlers/neptune_logger.py b/ignite/contrib/handlers/neptune_logger.py index 0d58305e0e9e..8830aa308ae5 100644 --- a/ignite/contrib/handlers/neptune_logger.py +++ b/ignite/contrib/handlers/neptune_logger.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.handlers.neptune_logger`` was moved to ``ignite.handlers.neptune_logger``. +"""``ignite.contrib.handlers.neptune_logger`` was moved to ``ignite.handlers.neptune_logger``. Note: ``ignite.contrib.handlers.neptune_logger`` was moved to ``ignite.handlers.neptune_logger``. Please refer to :mod:`~ignite.handlers.neptune_logger`. diff --git a/ignite/contrib/handlers/param_scheduler.py b/ignite/contrib/handlers/param_scheduler.py index 8a5643a6c09b..54e80a6f3ff1 100644 --- a/ignite/contrib/handlers/param_scheduler.py +++ b/ignite/contrib/handlers/param_scheduler.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.handlers.param_scheduler`` was moved to ``ignite.handlers.param_scheduler``. +"""``ignite.contrib.handlers.param_scheduler`` was moved to ``ignite.handlers.param_scheduler``. Note: ``ignite.contrib.handlers.param_scheduler`` was moved to ``ignite.handlers.param_scheduler``. Please refer to :mod:`~ignite.handlers.param_scheduler`. diff --git a/ignite/contrib/handlers/polyaxon_logger.py b/ignite/contrib/handlers/polyaxon_logger.py index bd2d82513277..54fbbe3c759e 100644 --- a/ignite/contrib/handlers/polyaxon_logger.py +++ b/ignite/contrib/handlers/polyaxon_logger.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.handlers.polyaxon_logger`` was moved to ``ignite.handlers.polyaxon_logger``. +"""``ignite.contrib.handlers.polyaxon_logger`` was moved to ``ignite.handlers.polyaxon_logger``. Note: ``ignite.contrib.handlers.polyaxon_logger`` was moved to ``ignite.handlers.polyaxon_logger``. Please refer to :mod:`~ignite.handlers.polyaxon_logger`. diff --git a/ignite/contrib/handlers/tensorboard_logger.py b/ignite/contrib/handlers/tensorboard_logger.py index 39e88f170c75..12936305eeec 100644 --- a/ignite/contrib/handlers/tensorboard_logger.py +++ b/ignite/contrib/handlers/tensorboard_logger.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.handlers.tensorboard_logger`` was moved to ``ignite.handlers.tensorboard_logger``. +"""``ignite.contrib.handlers.tensorboard_logger`` was moved to ``ignite.handlers.tensorboard_logger``. Note: ``ignite.contrib.handlers.tensorboard_logger`` was moved to ``ignite.handlers.tensorboard_logger``. Please refer to :mod:`~ignite.handlers.tensorboard_logger`. diff --git a/ignite/contrib/handlers/time_profilers.py b/ignite/contrib/handlers/time_profilers.py index 376499703448..f506211ddbfd 100644 --- a/ignite/contrib/handlers/time_profilers.py +++ b/ignite/contrib/handlers/time_profilers.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.handlers.time_profilers.py`` was moved to ``ignite.handlers.time_profilers``. +"""``ignite.contrib.handlers.time_profilers.py`` was moved to ``ignite.handlers.time_profilers``. Note: ``ignite.contrib.handlers.time_profilers`` was moved to ``ignite.handlers.time_profilers``. Please refer to :mod:`~ignite.handlers.time_profilers`. diff --git a/ignite/contrib/handlers/tqdm_logger.py b/ignite/contrib/handlers/tqdm_logger.py index 60393609b9e8..9f1f413bafe6 100644 --- a/ignite/contrib/handlers/tqdm_logger.py +++ b/ignite/contrib/handlers/tqdm_logger.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.handlers.tqdm_logger`` was moved to ``ignite.handlers.tqdm_logger``. +"""``ignite.contrib.handlers.tqdm_logger`` was moved to ``ignite.handlers.tqdm_logger``. Note: ``ignite.contrib.handlers.tqdm_logger`` was moved to ``ignite.handlers.tqdm_logger``. Please refer to :mod:`~ignite.handlers.tqdm_logger`. diff --git a/ignite/contrib/handlers/visdom_logger.py b/ignite/contrib/handlers/visdom_logger.py index f0eaf98530ce..de55b18857a0 100644 --- a/ignite/contrib/handlers/visdom_logger.py +++ b/ignite/contrib/handlers/visdom_logger.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.handlers.visdom_logger`` was moved to ``ignite.handlers.visdom_logger``. +"""``ignite.contrib.handlers.visdom_logger`` was moved to ``ignite.handlers.visdom_logger``. Note: ``ignite.contrib.handlers.visdom_logger`` was moved to ``ignite.handlers.visdom_logger``. Please refer to :mod:`~ignite.handlers.visdom_logger`. diff --git a/ignite/contrib/handlers/wandb_logger.py b/ignite/contrib/handlers/wandb_logger.py index f539db5b979b..88338e4209c9 100644 --- a/ignite/contrib/handlers/wandb_logger.py +++ b/ignite/contrib/handlers/wandb_logger.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.handlers.wandb_logger`` was moved to ``ignite.handlers.wandb_logger``. +"""``ignite.contrib.handlers.wandb_logger`` was moved to ``ignite.handlers.wandb_logger``. Note: ``ignite.contrib.handlers.wandb_logger`` was moved to ``ignite.handlers.wandb_logger``. Please refer to :mod:`~ignite.handlers.wandb_logger`. diff --git a/ignite/contrib/metrics/average_precision.py b/ignite/contrib/metrics/average_precision.py index 2940a277644c..cd039398b329 100644 --- a/ignite/contrib/metrics/average_precision.py +++ b/ignite/contrib/metrics/average_precision.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.average_precision`` was moved to ``ignite.metrics.average_precision``. +"""``ignite.contrib.metrics.average_precision`` was moved to ``ignite.metrics.average_precision``. Note: ``ignite.contrib.metrics.average_precision`` was moved to ``ignite.metrics.average_precision``. Please refer to :mod:`~ignite.metrics.average_precision`. diff --git a/ignite/contrib/metrics/cohen_kappa.py b/ignite/contrib/metrics/cohen_kappa.py index 7b99eb051329..164b4cf82291 100644 --- a/ignite/contrib/metrics/cohen_kappa.py +++ b/ignite/contrib/metrics/cohen_kappa.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.cohen_kappa`` was moved to ``ignite.metrics.cohen_kappa``. +"""``ignite.contrib.metrics.cohen_kappa`` was moved to ``ignite.metrics.cohen_kappa``. Note: ``ignite.contrib.metrics.cohen_kappa`` was moved to ``ignite.metrics.cohen_kappa``. Please refer to :mod:`~ignite.metrics.cohen_kappa`. diff --git a/ignite/contrib/metrics/gpu_info.py b/ignite/contrib/metrics/gpu_info.py index 07cdeca29e5f..992972326ad4 100644 --- a/ignite/contrib/metrics/gpu_info.py +++ b/ignite/contrib/metrics/gpu_info.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.gpu_info`` was moved to ``ignite.metrics.gpu_info``. +"""``ignite.contrib.metrics.gpu_info`` was moved to ``ignite.metrics.gpu_info``. Note: ``ignite.contrib.metrics.gpu_info`` was moved to ``ignite.metrics.gpu_info``. Please refer to :mod:`~ignite.metrics.gpu_info`. diff --git a/ignite/contrib/metrics/precision_recall_curve.py b/ignite/contrib/metrics/precision_recall_curve.py index c384aa8ebe84..9fc651f9b63d 100644 --- a/ignite/contrib/metrics/precision_recall_curve.py +++ b/ignite/contrib/metrics/precision_recall_curve.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.precision_recall_curve`` was moved to ``ignite.metrics.precision_recall_curve``. +"""``ignite.contrib.metrics.precision_recall_curve`` was moved to ``ignite.metrics.precision_recall_curve``. Note: ``ignite.contrib.metrics.precision_recall_curve`` was moved to ``ignite.metrics.precision_recall_curve``. Please refer to :mod:`~ignite.metrics.precision_recall_curve`. diff --git a/ignite/contrib/metrics/regression/canberra_metric.py b/ignite/contrib/metrics/regression/canberra_metric.py index 24195f174f2f..dcc608130e34 100644 --- a/ignite/contrib/metrics/regression/canberra_metric.py +++ b/ignite/contrib/metrics/regression/canberra_metric.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.canberra_metric`` was moved to ``ignite.metrics.regression.canberra_metric``. # noqa +"""``ignite.contrib.metrics.regression.canberra_metric`` was moved to ``ignite.metrics.regression.canberra_metric``. # noqa Note: ``ignite.contrib.metrics.regression.canberra_metric`` was moved to ``ignite.metrics.regression.canberra_metric``. # noqa Please refer to :mod:`~ignite.metrics.regression.canberra_metric`. diff --git a/ignite/contrib/metrics/regression/fractional_absolute_error.py b/ignite/contrib/metrics/regression/fractional_absolute_error.py index 85572af281a6..a257d2210613 100644 --- a/ignite/contrib/metrics/regression/fractional_absolute_error.py +++ b/ignite/contrib/metrics/regression/fractional_absolute_error.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.fractional_absolute_error`` was moved to ``ignite.metrics.regression.fractional_absolute_error``. # noqa +"""``ignite.contrib.metrics.regression.fractional_absolute_error`` was moved to ``ignite.metrics.regression.fractional_absolute_error``. # noqa Note: ``ignite.contrib.metrics.regression.fractional_absolute_error`` was moved to ``ignite.metrics.regression.fractional_absolute_error``. # noqa Please refer to :mod:`~ignite.metrics.regression.fractional_absolute_error`. diff --git a/ignite/contrib/metrics/regression/fractional_bias.py b/ignite/contrib/metrics/regression/fractional_bias.py index a0749979da67..236d2e60c8ed 100644 --- a/ignite/contrib/metrics/regression/fractional_bias.py +++ b/ignite/contrib/metrics/regression/fractional_bias.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.fractional_bias`` was moved to ``ignite.metrics.regression.fractional_bias``. # noqa +"""``ignite.contrib.metrics.regression.fractional_bias`` was moved to ``ignite.metrics.regression.fractional_bias``. # noqa Note: ``ignite.contrib.metrics.regression.fractional_bias`` was moved to ``ignite.metrics.regression.fractional_bias``. # noqa Please refer to :mod:`~ignite.metrics.regression.fractional_bias`. diff --git a/ignite/contrib/metrics/regression/geometric_mean_absolute_error.py b/ignite/contrib/metrics/regression/geometric_mean_absolute_error.py index 49308cd1a02f..9b5be64778c8 100644 --- a/ignite/contrib/metrics/regression/geometric_mean_absolute_error.py +++ b/ignite/contrib/metrics/regression/geometric_mean_absolute_error.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.geometric_mean_absolute_error`` was moved to ``ignite.metrics.regression.geometric_mean_absolute_error``. # noqa +"""``ignite.contrib.metrics.regression.geometric_mean_absolute_error`` was moved to ``ignite.metrics.regression.geometric_mean_absolute_error``. # noqa Note: ``ignite.contrib.metrics.regression.geometric_mean_absolute_error`` was moved to ``ignite.metrics.regression.geometric_mean_absolute_error``. # noqa Please refer to :mod:`~ignite.metrics.regression.geometric_mean_absolute_error`. diff --git a/ignite/contrib/metrics/regression/geometric_mean_relative_absolute_error.py b/ignite/contrib/metrics/regression/geometric_mean_relative_absolute_error.py index 992f7dbe4e5c..7c1af66bdd4a 100644 --- a/ignite/contrib/metrics/regression/geometric_mean_relative_absolute_error.py +++ b/ignite/contrib/metrics/regression/geometric_mean_relative_absolute_error.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.geometric_mean_relative_absolute_error`` was moved to ``ignite.metrics.regression.geometric_mean_relative_absolute_error``. # noqa +"""``ignite.contrib.metrics.regression.geometric_mean_relative_absolute_error`` was moved to ``ignite.metrics.regression.geometric_mean_relative_absolute_error``. # noqa Note: ``ignite.contrib.metrics.regression.geometric_mean_relative_absolute_error`` was moved to ``ignite.metrics.regression.geometric_mean_relative_absolute_error``. # noqa Please refer to :mod:`~ignite.metrics.regression.geometric_mean_relative_absolute_error`. diff --git a/ignite/contrib/metrics/regression/manhattan_distance.py b/ignite/contrib/metrics/regression/manhattan_distance.py index 912596c8ab22..0f13413723bb 100644 --- a/ignite/contrib/metrics/regression/manhattan_distance.py +++ b/ignite/contrib/metrics/regression/manhattan_distance.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.manhattan_distance`` was moved to ``ignite.metrics.regression.manhattan_distance``. # noqa +"""``ignite.contrib.metrics.regression.manhattan_distance`` was moved to ``ignite.metrics.regression.manhattan_distance``. # noqa Note: ``ignite.contrib.metrics.regression.manhattan_distance`` was moved to ``ignite.metrics.regression.manhattan_distance``. # noqa Please refer to :mod:`~ignite.metrics.regression.manhattan_distance`. diff --git a/ignite/contrib/metrics/regression/maximum_absolute_error.py b/ignite/contrib/metrics/regression/maximum_absolute_error.py index 7e71dc9a41d7..7b974d274a03 100644 --- a/ignite/contrib/metrics/regression/maximum_absolute_error.py +++ b/ignite/contrib/metrics/regression/maximum_absolute_error.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.maximum_absolute_error`` was moved to ``ignite.metrics.regression.maximum_absolute_error``. # noqa +"""``ignite.contrib.metrics.regression.maximum_absolute_error`` was moved to ``ignite.metrics.regression.maximum_absolute_error``. # noqa Note: ``ignite.contrib.metrics.regression.maximum_absolute_error`` was moved to ``ignite.metrics.regression.maximum_absolute_error``. # noqa Please refer to :mod:`~ignite.metrics.regression.maximum_absolute_error`. diff --git a/ignite/contrib/metrics/regression/mean_absolute_relative_error.py b/ignite/contrib/metrics/regression/mean_absolute_relative_error.py index 6e2fbd2df023..c57b2bcba679 100644 --- a/ignite/contrib/metrics/regression/mean_absolute_relative_error.py +++ b/ignite/contrib/metrics/regression/mean_absolute_relative_error.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.mean_absolute_relative_error`` was moved to ``ignite.metrics.regression.mean_absolute_relative_error``. # noqa +"""``ignite.contrib.metrics.regression.mean_absolute_relative_error`` was moved to ``ignite.metrics.regression.mean_absolute_relative_error``. # noqa Note: ``ignite.contrib.metrics.regression.mean_absolute_relative_error`` was moved to ``ignite.metrics.regression.mean_absolute_relative_error``. # noqa Please refer to :mod:`~ignite.metrics.regression.mean_absolute_relative_error`. diff --git a/ignite/contrib/metrics/regression/mean_error.py b/ignite/contrib/metrics/regression/mean_error.py index 1a4dbdd1ad6d..b2d6251409fb 100644 --- a/ignite/contrib/metrics/regression/mean_error.py +++ b/ignite/contrib/metrics/regression/mean_error.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.mean_error`` was moved to ``ignite.metrics.regression.mean_error``. # noqa +"""``ignite.contrib.metrics.regression.mean_error`` was moved to ``ignite.metrics.regression.mean_error``. # noqa Note: ``ignite.contrib.metrics.regression.mean_error`` was moved to ``ignite.metrics.regression.mean_error``. # noqa Please refer to :mod:`~ignite.metrics.regression.mean_error`. diff --git a/ignite/contrib/metrics/regression/mean_normalized_bias.py b/ignite/contrib/metrics/regression/mean_normalized_bias.py index 0a3523555cac..3c57598a724f 100644 --- a/ignite/contrib/metrics/regression/mean_normalized_bias.py +++ b/ignite/contrib/metrics/regression/mean_normalized_bias.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.mean_normalized_bias`` was moved to ``ignite.metrics.regression.mean_normalized_bias``. # noqa +"""``ignite.contrib.metrics.regression.mean_normalized_bias`` was moved to ``ignite.metrics.regression.mean_normalized_bias``. # noqa Note: ``ignite.contrib.metrics.regression.mean_normalized_bias`` was moved to ``ignite.metrics.regression.mean_normalized_bias``. # noqa Please refer to :mod:`~ignite.metrics.regression.mean_normalized_bias`. diff --git a/ignite/contrib/metrics/regression/median_absolute_error.py b/ignite/contrib/metrics/regression/median_absolute_error.py index 98b0f1438f12..06af5a065eec 100644 --- a/ignite/contrib/metrics/regression/median_absolute_error.py +++ b/ignite/contrib/metrics/regression/median_absolute_error.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.median_absolute_error`` was moved to ``ignite.metrics.regression.median_absolute_error``. # noqa +"""``ignite.contrib.metrics.regression.median_absolute_error`` was moved to ``ignite.metrics.regression.median_absolute_error``. # noqa Note: ``ignite.contrib.metrics.regression.median_absolute_error`` was moved to ``ignite.metrics.regression.median_absolute_error``. # noqa Please refer to :mod:`~ignite.metrics.regression.median_absolute_error`. diff --git a/ignite/contrib/metrics/regression/median_absolute_percentage_error.py b/ignite/contrib/metrics/regression/median_absolute_percentage_error.py index cd74e9e74953..94ed5f2b014c 100644 --- a/ignite/contrib/metrics/regression/median_absolute_percentage_error.py +++ b/ignite/contrib/metrics/regression/median_absolute_percentage_error.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.median_absolute_percentage_error`` was moved to ``ignite.metrics.regression.median_absolute_percentage_error``. # noqa +"""``ignite.contrib.metrics.regression.median_absolute_percentage_error`` was moved to ``ignite.metrics.regression.median_absolute_percentage_error``. # noqa Note: ``ignite.contrib.metrics.regression.median_absolute_percentage_error`` was moved to ``ignite.metrics.regression.median_absolute_percentage_error``. # noqa Please refer to :mod:`~ignite.metrics.regression.median_absolute_percentage_error`. diff --git a/ignite/contrib/metrics/regression/median_relative_absolute_error.py b/ignite/contrib/metrics/regression/median_relative_absolute_error.py index 56769e78a205..8dc1effdd139 100644 --- a/ignite/contrib/metrics/regression/median_relative_absolute_error.py +++ b/ignite/contrib/metrics/regression/median_relative_absolute_error.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.median_relative_absolute_error`` was moved to ``ignite.metrics.regression.median_relative_absolute_error``. # noqa +"""``ignite.contrib.metrics.regression.median_relative_absolute_error`` was moved to ``ignite.metrics.regression.median_relative_absolute_error``. # noqa Note: ``ignite.contrib.metrics.regression.median_relative_absolute_error`` was moved to ``ignite.metrics.regression.median_relative_absolute_error``. # noqa Please refer to :mod:`~ignite.metrics.regression.median_relative_absolute_error`. diff --git a/ignite/contrib/metrics/regression/r2_score.py b/ignite/contrib/metrics/regression/r2_score.py index 99bdbaa6fb45..cc442d430134 100644 --- a/ignite/contrib/metrics/regression/r2_score.py +++ b/ignite/contrib/metrics/regression/r2_score.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.r2_score`` was moved to ``ignite.metrics.regression.r2_score``. # noqa +"""``ignite.contrib.metrics.regression.r2_score`` was moved to ``ignite.metrics.regression.r2_score``. # noqa Note: ``ignite.contrib.metrics.regression.r2_score`` was moved to ``ignite.metrics.regression.r2_score``. # noqa Please refer to :mod:`~ignite.metrics.regression.r2_score`. diff --git a/ignite/contrib/metrics/regression/wave_hedges_distance.py b/ignite/contrib/metrics/regression/wave_hedges_distance.py index fb3ccf0d7b06..c0f6ed2ea495 100644 --- a/ignite/contrib/metrics/regression/wave_hedges_distance.py +++ b/ignite/contrib/metrics/regression/wave_hedges_distance.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.regression.wave_hedges_distance`` was moved to ``ignite.metrics.regression.wave_hedges_distance``. # noqa +"""``ignite.contrib.metrics.regression.wave_hedges_distance`` was moved to ``ignite.metrics.regression.wave_hedges_distance``. # noqa Note: ``ignite.contrib.metrics.regression.wave_hedges_distance`` was moved to ``ignite.metrics.regression.wave_hedges_distance``. # noqa Please refer to :mod:`~ignite.metrics.regression.wave_hedges_distance`. diff --git a/ignite/contrib/metrics/roc_auc.py b/ignite/contrib/metrics/roc_auc.py index a8cd774b3dbc..0190f8afc27b 100644 --- a/ignite/contrib/metrics/roc_auc.py +++ b/ignite/contrib/metrics/roc_auc.py @@ -1,4 +1,4 @@ -""" ``ignite.contrib.metrics.roc_auc`` was moved to ``ignite.metrics.roc_auc``. +"""``ignite.contrib.metrics.roc_auc`` was moved to ``ignite.metrics.roc_auc``. Note: ``ignite.contrib.metrics.roc_auc`` was moved to ``ignite.metrics.roc_auc``. Please refer to :mod:`~ignite.metrics.roc_auc`. diff --git a/ignite/distributed/comp_models/__init__.py b/ignite/distributed/comp_models/__init__.py index 9b7110d15cb6..36f44363ee72 100644 --- a/ignite/distributed/comp_models/__init__.py +++ b/ignite/distributed/comp_models/__init__.py @@ -13,9 +13,9 @@ from ignite.distributed.comp_models.xla import _XlaDistModel -def setup_available_computation_models() -> ( - tuple[type[_SerialModel | _NativeDistModel | _XlaDistModel | _HorovodDistModel], ...] -): +def setup_available_computation_models() -> tuple[ + type[_SerialModel | _NativeDistModel | _XlaDistModel | _HorovodDistModel], ... +]: models: list[type[_SerialModel | _NativeDistModel | _XlaDistModel | _HorovodDistModel]] = [ _SerialModel, ] diff --git a/ignite/engine/deterministic.py b/ignite/engine/deterministic.py index f6674b58cf04..e16c72b12db3 100644 --- a/ignite/engine/deterministic.py +++ b/ignite/engine/deterministic.py @@ -228,7 +228,8 @@ def _setup_engine(self) -> None: batch_sampler = self.state.dataloader.batch_sampler if not (batch_sampler is None or isinstance(batch_sampler, ReproducibleBatchSampler)): self.state.dataloader = update_dataloader( - self.state.dataloader, ReproducibleBatchSampler(batch_sampler) # type: ignore[arg-type] + self.state.dataloader, + ReproducibleBatchSampler(batch_sampler), # type: ignore[arg-type] ) iteration = self.state.iteration diff --git a/ignite/engine/engine.py b/ignite/engine/engine.py index cdafecc4ae55..72cf343871bc 100644 --- a/ignite/engine/engine.py +++ b/ignite/engine/engine.py @@ -1081,8 +1081,7 @@ def _run_once_on_dataset_as_gen(self) -> Generator[State, None, float]: try: if self._dataloader_iter is None: raise RuntimeError( - "Internal error, self._dataloader_iter is None. " - "Please, file an issue if you encounter this error." + "Internal error, self._dataloader_iter is None. Please, file an issue if you encounter this error." ) while True: @@ -1269,8 +1268,7 @@ def _run_once_on_dataset_legacy(self) -> float: try: if self._dataloader_iter is None: raise RuntimeError( - "Internal error, self._dataloader_iter is None. " - "Please, file an issue if you encounter this error." + "Internal error, self._dataloader_iter is None. Please, file an issue if you encounter this error." ) while True: diff --git a/ignite/handlers/ema_handler.py b/ignite/handlers/ema_handler.py index 75d83146a21c..df96a82c705d 100644 --- a/ignite/handlers/ema_handler.py +++ b/ignite/handlers/ema_handler.py @@ -165,8 +165,7 @@ def __init__( if not isinstance(model, nn.Module): raise ValueError( - f"model should be an instance of nn.Module or its subclasses, but got" - f"model: {model.__class__.__name__}" + f"model should be an instance of nn.Module or its subclasses, but gotmodel: {model.__class__.__name__}" ) if isinstance(model, nn.parallel.DistributedDataParallel): @@ -179,7 +178,7 @@ def __init__( if handle_buffers not in ("copy", "update", "ema_train"): raise ValueError( - f"handle_buffers can only be one of 'copy', 'update', 'ema_train', " f"but got {handle_buffers}" + f"handle_buffers can only be one of 'copy', 'update', 'ema_train', but got {handle_buffers}" ) self.handle_buffers = handle_buffers diff --git a/ignite/handlers/mlflow_logger.py b/ignite/handlers/mlflow_logger.py index c253f6f1edd6..751b9c5ae43f 100644 --- a/ignite/handlers/mlflow_logger.py +++ b/ignite/handlers/mlflow_logger.py @@ -282,8 +282,7 @@ def __call__(self, engine: Engine, logger: MLflowLogger, event_name: str | Event if not isinstance(global_step, int): raise TypeError( - f"global_step must be int, got {type(global_step)}." - " Please check the output of global_step_transform." + f"global_step must be int, got {type(global_step)}. Please check the output of global_step_transform." ) # Additionally recheck metric names as MLflow rejects non-valid names with MLflowException diff --git a/ignite/handlers/neptune_logger.py b/ignite/handlers/neptune_logger.py index fcbc9e899f73..5e77ce951381 100644 --- a/ignite/handlers/neptune_logger.py +++ b/ignite/handlers/neptune_logger.py @@ -334,8 +334,7 @@ def __call__(self, engine: Engine, logger: NeptuneLogger, event_name: str | Even if not isinstance(global_step, int): raise TypeError( - f"global_step must be int, got {type(global_step)}." - " Please check the output of global_step_transform." + f"global_step must be int, got {type(global_step)}. Please check the output of global_step_transform." ) for key, value in metrics.items(): diff --git a/ignite/handlers/param_scheduler.py b/ignite/handlers/param_scheduler.py index e7373c781388..0e780de713a5 100644 --- a/ignite/handlers/param_scheduler.py +++ b/ignite/handlers/param_scheduler.py @@ -685,7 +685,7 @@ def __init__(self, schedulers: list[ParamScheduler], durations: list[int], save_ if len(schedulers) != len(durations) + 1: raise ValueError( - "Incorrect number schedulers or duration values, " f"given {len(schedulers)} and {len(durations)}" + f"Incorrect number schedulers or duration values, given {len(schedulers)} and {len(durations)}" ) for i, scheduler in enumerate(schedulers): @@ -1727,8 +1727,7 @@ def simulate_values( # type: ignore[override] """ if len(metric_values) != num_events: raise ValueError( - "Length of argument metric_values should be equal to num_events. " - f"{len(metric_values)} != {num_events}" + f"Length of argument metric_values should be equal to num_events. {len(metric_values)} != {num_events}" ) keys_to_remove = ["optimizer", "metric_name", "save_history"] diff --git a/ignite/handlers/time_profilers.py b/ignite/handlers/time_profilers.py index 5a8adb92996c..9798f0c61346 100644 --- a/ignite/handlers/time_profilers.py +++ b/ignite/handlers/time_profilers.py @@ -421,8 +421,7 @@ def odict_to_str(d: Mapping) -> str: others.update(results["event_handlers_names"]) - output_message: str = ( - """ + output_message: str = """ ---------------------------------------------------- | Time profiling stats (in seconds): | ---------------------------------------------------- @@ -455,10 +454,9 @@ def odict_to_str(d: Mapping) -> str: - Events.COMPLETED: {COMPLETED_names} {COMPLETED} """.format( - processing_stats=odict_to_str(results["processing_stats"]), - dataflow_stats=odict_to_str(results["dataflow_stats"]), - **others, - ) + processing_stats=odict_to_str(results["processing_stats"]), + dataflow_stats=odict_to_str(results["dataflow_stats"]), + **others, ) print(output_message) return output_message diff --git a/ignite/handlers/tqdm_logger.py b/ignite/handlers/tqdm_logger.py index 4125bd08503f..db56c3c27c4b 100644 --- a/ignite/handlers/tqdm_logger.py +++ b/ignite/handlers/tqdm_logger.py @@ -1,5 +1,6 @@ # -*- coding: utf-8 -*- """TQDM logger.""" + from collections import OrderedDict from collections.abc import Callable from typing import Any @@ -128,8 +129,7 @@ def __init__( from tqdm.autonotebook import tqdm except ImportError: raise ModuleNotFoundError( - "This contrib module requires tqdm to be installed. " - "Please install it with command: \n pip install tqdm" + "This contrib module requires tqdm to be installed. Please install it with command: \n pip install tqdm" ) self.pbar_cls = tqdm diff --git a/ignite/metrics/nlp/bleu.py b/ignite/metrics/nlp/bleu.py index 621786ef2884..9f7421418ba3 100644 --- a/ignite/metrics/nlp/bleu.py +++ b/ignite/metrics/nlp/bleu.py @@ -165,8 +165,7 @@ def _n_gram_counter( ) -> tuple[int, int]: if len(references) != len(candidates): raise ValueError( - f"nb of candidates should be equal to nb of reference lists ({len(candidates)} != " - f"{len(references)})" + f"nb of candidates should be equal to nb of reference lists ({len(candidates)} != {len(references)})" ) hyp_lengths = 0 diff --git a/ignite/metrics/running_average.py b/ignite/metrics/running_average.py index 898d24f33a72..a2b5426c5cb7 100644 --- a/ignite/metrics/running_average.py +++ b/ignite/metrics/running_average.py @@ -124,8 +124,7 @@ def output_transform(x: Any) -> Any: else: if output_transform is None: raise ValueError( - "Argument output_transform should not be None if src corresponds " - "to the output of process function." + "Argument output_transform should not be None if src corresponds to the output of process function." ) self.src = None if device is None: diff --git a/ignite/utils.py b/ignite/utils.py index d504a2a9dea2..2875a2b5e5d2 100644 --- a/ignite/utils.py +++ b/ignite/utils.py @@ -177,7 +177,7 @@ class _CollectionItem: def __init__(self, collection: dict | list, key: int | str) -> None: if not isinstance(collection, (dict, list)): raise TypeError( - f"Input type is expected to be a mapping or list, but got {type(collection)} " f"for input key '{key}'." + f"Input type is expected to be a mapping or list, but got {type(collection)} for input key '{key}'." ) if isinstance(collection, list) and isinstance(key, str): raise ValueError("Key should be int for collection of type list") diff --git a/tests/ignite/conftest.py b/tests/ignite/conftest.py index a9a35de69813..27a7fe3ccc8e 100644 --- a/tests/ignite/conftest.py +++ b/tests/ignite/conftest.py @@ -249,7 +249,7 @@ def distributed_context_single_node_gloo(local_rank, world_size): temp_file = tempfile.NamedTemporaryFile(delete=False) # can't use backslashes in f-strings backslash = "\\" - init_method = f'file:///{temp_file.name.replace(backslash, "/")}' + init_method = f"file:///{temp_file.name.replace(backslash, '/')}" else: free_port = _setup_free_port(local_rank) init_method = f"tcp://localhost:{free_port}" @@ -470,7 +470,7 @@ def distributed(request, local_rank, world_size): temp_file = tempfile.NamedTemporaryFile(delete=False) # can't use backslashes in f-strings backslash = "\\" - init_method = f'file:///{temp_file.name.replace(backslash, "/")}' + init_method = f"file:///{temp_file.name.replace(backslash, '/')}" else: temp_file = None free_port = _setup_free_port(local_rank) diff --git a/tests/ignite/contrib/engines/test_common.py b/tests/ignite/contrib/engines/test_common.py index 95ab7e280faa..e4a65eb72014 100644 --- a/tests/ignite/contrib/engines/test_common.py +++ b/tests/ignite/contrib/engines/test_common.py @@ -126,9 +126,9 @@ def update_fn(engine, batch): assert any([v in c for c in checkpoints]) # Check LR scheduling - assert optimizer.param_groups[0]["lr"] <= lr * gamma ** ( - (num_iters * num_epochs - 1) // step_size - ), f"{optimizer.param_groups[0]['lr']} vs {lr * gamma ** ((num_iters * num_epochs - 1) // step_size)}" + assert optimizer.param_groups[0]["lr"] <= lr * gamma ** ((num_iters * num_epochs - 1) // step_size), ( + f"{optimizer.param_groups[0]['lr']} vs {lr * gamma ** ((num_iters * num_epochs - 1) // step_size)}" + ) def test_asserts_setup_common_training_handlers(): diff --git a/tests/ignite/distributed/comp_models/test_native.py b/tests/ignite/distributed/comp_models/test_native.py index 09e4d3054601..14bdccefa7c6 100644 --- a/tests/ignite/distributed/comp_models/test_native.py +++ b/tests/ignite/distributed/comp_models/test_native.py @@ -38,10 +38,7 @@ ), ( "node[4-8,12,16-20,22,24-26]", - "node4,node5,node6,node7,node8," - "node12,node16,node17,node18," - "node19,node20,node22,node24," - "node25,node26", + "node4,node5,node6,node7,node8,node12,node16,node17,node18,node19,node20,node22,node24,node25,node26", ), ("machine2-[02-4]vm1", "machine2-02vm1,machine2-03vm1,machine2-04vm1"), ( @@ -578,47 +575,78 @@ def test__native_dist_model_init_method_is_not_none(world_size, local_rank, get_ # fmt: off # usual SLURM env ( - { - "SLURM_PROCID": "1", "SLURM_LOCALID": "1", "SLURM_NTASKS": "2", "SLURM_JOB_NUM_NODES": "1", - "SLURM_JOB_NODELIST": "c1", "SLURM_JOB_ID": "12345", + "SLURM_PROCID": "1", + "SLURM_LOCALID": "1", + "SLURM_NTASKS": "2", + "SLURM_JOB_NUM_NODES": "1", + "SLURM_JOB_NODELIST": "c1", + "SLURM_JOB_ID": "12345", }, - [1, 1, 2, "c1", 17345] + [1, 1, 2, "c1", 17345], ), # usual SLURM env mnode ( { - "SLURM_PROCID": "5", "SLURM_LOCALID": "1", "SLURM_NTASKS": "8", "SLURM_JOB_NUM_NODES": "2", - "SLURM_JOB_NODELIST": "c1, c2", "SLURM_JOB_ID": "12345", + "SLURM_PROCID": "5", + "SLURM_LOCALID": "1", + "SLURM_NTASKS": "8", + "SLURM_JOB_NUM_NODES": "2", + "SLURM_JOB_NODELIST": "c1, c2", + "SLURM_JOB_ID": "12345", }, - [5, 1, 8, "c1", 17345] + [5, 1, 8, "c1", 17345], ), # usual SLURM env 1 node, 1 task + torch.distributed.launch ( { - "SLURM_PROCID": "0", "SLURM_LOCALID": "0", "SLURM_NTASKS": "1", "SLURM_JOB_NUM_NODES": "1", - "SLURM_JOB_NODELIST": "c1", "SLURM_JOB_ID": "12345", - "MASTER_ADDR": "127.0.0.1", "MASTER_PORT": "2233", "RANK": "2", "LOCAL_RANK": "2", "WORLD_SIZE": "8", + "SLURM_PROCID": "0", + "SLURM_LOCALID": "0", + "SLURM_NTASKS": "1", + "SLURM_JOB_NUM_NODES": "1", + "SLURM_JOB_NODELIST": "c1", + "SLURM_JOB_ID": "12345", + "MASTER_ADDR": "127.0.0.1", + "MASTER_PORT": "2233", + "RANK": "2", + "LOCAL_RANK": "2", + "WORLD_SIZE": "8", }, - [2, 2, 8, "127.0.0.1", 2233] + [2, 2, 8, "127.0.0.1", 2233], ), # usual SLURM env + enroot's pytorch hook ( { - "SLURM_PROCID": "3", "SLURM_LOCALID": "3", "SLURM_NTASKS": "4", "SLURM_JOB_NUM_NODES": "1", - "SLURM_JOB_NODELIST": "c1", "SLURM_JOB_ID": "12345", - "MASTER_ADDR": "c1", "MASTER_PORT": "12233", "RANK": "3", "LOCAL_RANK": "3", "WORLD_SIZE": "4", + "SLURM_PROCID": "3", + "SLURM_LOCALID": "3", + "SLURM_NTASKS": "4", + "SLURM_JOB_NUM_NODES": "1", + "SLURM_JOB_NODELIST": "c1", + "SLURM_JOB_ID": "12345", + "MASTER_ADDR": "c1", + "MASTER_PORT": "12233", + "RANK": "3", + "LOCAL_RANK": "3", + "WORLD_SIZE": "4", }, - [3, 3, 4, "c1", 12233] + [3, 3, 4, "c1", 12233], ), # usual SLURM env mnode + enroot's pytorch hook ( { - "SLURM_PROCID": "3", "SLURM_LOCALID": "1", "SLURM_NTASKS": "4", "SLURM_JOB_NUM_NODES": "2", - "SLURM_JOB_NODELIST": "c1, c2", "SLURM_JOB_ID": "12345", - "MASTER_ADDR": "c1", "MASTER_PORT": "12233", "RANK": "3", "LOCAL_RANK": "1", "WORLD_SIZE": "4" + "SLURM_PROCID": "3", + "SLURM_LOCALID": "1", + "SLURM_NTASKS": "4", + "SLURM_JOB_NUM_NODES": "2", + "SLURM_JOB_NODELIST": "c1, c2", + "SLURM_JOB_ID": "12345", + "MASTER_ADDR": "c1", + "MASTER_PORT": "12233", + "RANK": "3", + "LOCAL_RANK": "1", + "WORLD_SIZE": "4", }, - [3, 1, 4, "c1", 12233] + [3, 1, 4, "c1", 12233], ), # fmt: on ], diff --git a/tests/ignite/distributed/test_auto.py b/tests/ignite/distributed/test_auto.py index bfe680dc4f97..b26c4a7417d3 100644 --- a/tests/ignite/distributed/test_auto.py +++ b/tests/ignite/distributed/test_auto.py @@ -145,9 +145,9 @@ def _test_auto_model(model, ws, device, sync_bn=False, **kwargs): else: assert isinstance(model, nn.Module) - assert all( - [p.device.type == torch.device(device).type for p in model.parameters()] - ), f"{[p.device.type for p in model.parameters()]} vs {torch.device(device).type}" + assert all([p.device.type == torch.device(device).type for p in model.parameters()]), ( + f"{[p.device.type for p in model.parameters()]} vs {torch.device(device).type}" + ) def _test_auto_model_optimizer(ws, device): @@ -303,9 +303,9 @@ def test_dist_proxy_sampler(): set_indices_per_rank = set(indices_per_rank) set_true_indices = set(true_indices) - assert ( - set_indices_per_rank == set_true_indices - ), f"{set_true_indices - set_indices_per_rank} | {set_indices_per_rank - set_true_indices}" + assert set_indices_per_rank == set_true_indices, ( + f"{set_true_indices - set_indices_per_rank} | {set_indices_per_rank - set_true_indices}" + ) with pytest.raises(TypeError, match=r"Argument sampler should be instance of torch Sampler"): DistributedProxySampler(None) diff --git a/tests/ignite/engine/test_deterministic.py b/tests/ignite/engine/test_deterministic.py index 9af116956221..80fae5a7bf0c 100644 --- a/tests/ignite/engine/test_deterministic.py +++ b/tests/ignite/engine/test_deterministic.py @@ -185,9 +185,9 @@ def _test(epoch_length=None): batch_checker = BatchChecker(data, init_counter=resume_iteration) def update_fn(_, batch): - assert batch_checker.check( - batch - ), f"{resume_iteration} | {batch_checker.counter}: {batch_checker.true_batch} vs {batch}" + assert batch_checker.check(batch), ( + f"{resume_iteration} | {batch_checker.counter}: {batch_checker.true_batch} vs {batch}" + ) engine = DeterministicEngine(update_fn) @@ -221,9 +221,9 @@ def _test(epoch_length=None): batch_checker = BatchChecker(data, init_counter=resume_epoch * epoch_length) def update_fn(_, batch): - assert batch_checker.check( - batch - ), f"{resume_epoch} | {batch_checker.counter}: {batch_checker.true_batch} vs {batch}" + assert batch_checker.check(batch), ( + f"{resume_epoch} | {batch_checker.counter}: {batch_checker.true_batch} vs {batch}" + ) engine = DeterministicEngine(update_fn) @@ -297,9 +297,9 @@ def _(engine): def update_fn(_, batch): batch_to_device = batch.to(device) - assert batch_checker.check( - batch - ), f"{num_workers} {resume_epoch} | {batch_checker.counter}: {batch_checker.true_batch} vs {batch}" + assert batch_checker.check(batch), ( + f"{num_workers} {resume_epoch} | {batch_checker.counter}: {batch_checker.true_batch} vs {batch}" + ) engine = DeterministicEngine(update_fn) @@ -401,9 +401,9 @@ def _(engine): def update_fn(_, batch): batch_to_device = batch.to(device) cfg_msg = f"{num_workers} {resume_iteration}" - assert batch_checker.check( - batch - ), f"{cfg_msg} | {batch_checker.counter}: {batch_checker.true_batch} vs {batch}" + assert batch_checker.check(batch), ( + f"{cfg_msg} | {batch_checker.counter}: {batch_checker.true_batch} vs {batch}" + ) engine = DeterministicEngine(update_fn) @@ -420,9 +420,9 @@ def _(engine): torch.manual_seed(12) engine.run(resume_dataloader) assert engine.state.epoch == max_epochs - assert ( - engine.state.iteration == epoch_length * max_epochs - ), f"{num_workers}, {resume_iteration} | {engine.state.iteration} vs {epoch_length * max_epochs}" + assert engine.state.iteration == epoch_length * max_epochs, ( + f"{num_workers}, {resume_iteration} | {engine.state.iteration} vs {epoch_length * max_epochs}" + ) _test() if sampler_type != "distributed": @@ -467,9 +467,9 @@ def update_fn(_, batch): batch_checker = BatchChecker(seen_batchs, init_counter=resume_epoch * epoch_length) def update_fn(_, batch): - assert batch_checker.check( - batch - ), f"{resume_epoch} | {batch_checker.counter}: {batch_checker.true_batch} vs {batch}" + assert batch_checker.check(batch), ( + f"{resume_epoch} | {batch_checker.counter}: {batch_checker.true_batch} vs {batch}" + ) engine = DeterministicEngine(update_fn) @@ -520,9 +520,9 @@ def update_fn(_, batch): batch_checker = BatchChecker(seen_batchs, init_counter=resume_iteration) def update_fn(_, batch): - assert batch_checker.check( - batch - ), f"{resume_iteration} | {batch_checker.counter}: {batch_checker.true_batch} vs {batch}" + assert batch_checker.check(batch), ( + f"{resume_iteration} | {batch_checker.counter}: {batch_checker.true_batch} vs {batch}" + ) engine = DeterministicEngine(update_fn) @@ -533,9 +533,9 @@ def update_fn(_, batch): torch.manual_seed(24) engine.run(infinite_data_iterator()) assert engine.state.epoch == max_epochs - assert ( - engine.state.iteration == epoch_length * max_epochs - ), f"{resume_iteration} | {engine.state.iteration} vs {epoch_length * max_epochs}" + assert engine.state.iteration == epoch_length * max_epochs, ( + f"{resume_iteration} | {engine.state.iteration} vs {epoch_length * max_epochs}" + ) _test() _test(50) diff --git a/tests/ignite/engine/test_memory_leaks.py b/tests/ignite/engine/test_memory_leaks.py index e5a52fca4f7d..9100069468bc 100644 --- a/tests/ignite/engine/test_memory_leaks.py +++ b/tests/ignite/engine/test_memory_leaks.py @@ -28,7 +28,6 @@ def test_memory_leak(self, with_handler): counter = 0 class EngineForTests(Engine): - def __del__(self): nonlocal counter counter += 1 diff --git a/tests/ignite/handlers/test_neptune_logger.py b/tests/ignite/handlers/test_neptune_logger.py index 02b02fa828c1..67b1e058a5a5 100644 --- a/tests/ignite/handlers/test_neptune_logger.py +++ b/tests/ignite/handlers/test_neptune_logger.py @@ -200,13 +200,13 @@ def test_output_handler_metric_names(): wrapper(mock_engine, logger, Events.ITERATION_STARTED) - assert_logger_called_once_with(logger, "tag/a", 123), - assert_logger_called_once_with(logger, "tag/b/c/0", 2.34), - assert_logger_called_once_with(logger, "tag/b/c/1/d", 1), - assert_logger_called_once_with(logger, "tag/c/0", 22), - assert_logger_called_once_with(logger, "tag/c/1/0", 33), - assert_logger_called_once_with(logger, "tag/c/1/1", -5.5), - assert_logger_called_once_with(logger, "tag/c/2/e", 32.1), + (assert_logger_called_once_with(logger, "tag/a", 123),) + (assert_logger_called_once_with(logger, "tag/b/c/0", 2.34),) + (assert_logger_called_once_with(logger, "tag/b/c/1/d", 1),) + (assert_logger_called_once_with(logger, "tag/c/0", 22),) + (assert_logger_called_once_with(logger, "tag/c/1/0", 33),) + (assert_logger_called_once_with(logger, "tag/c/1/1", -5.5),) + (assert_logger_called_once_with(logger, "tag/c/2/e", 32.1),) logger.stop() diff --git a/tests/ignite/handlers/test_param_scheduler.py b/tests/ignite/handlers/test_param_scheduler.py index 6cc3b2893f6d..19f6afc1d28b 100644 --- a/tests/ignite/handlers/test_param_scheduler.py +++ b/tests/ignite/handlers/test_param_scheduler.py @@ -1070,12 +1070,12 @@ def save_lr(engine): assert lrs == pytest.approx([v for _, v in simulated_values]) assert lrs[0] == pytest.approx(warmup_start_value), f"lrs={lrs[: warmup_duration + num_iterations]}" - assert lrs[warmup_duration - 1] == pytest.approx( - expected_warmup_end_value - ), f"lrs={lrs[: warmup_duration + num_iterations]}" - assert lrs[warmup_duration] == pytest.approx( - warmup_end_next_value - ), f"lrs={lrs[: warmup_duration + num_iterations]}" + assert lrs[warmup_duration - 1] == pytest.approx(expected_warmup_end_value), ( + f"lrs={lrs[: warmup_duration + num_iterations]}" + ) + assert lrs[warmup_duration] == pytest.approx(warmup_end_next_value), ( + f"lrs={lrs[: warmup_duration + num_iterations]}" + ) scheduler.load_state_dict(state_dict) diff --git a/tests/ignite/metrics/nlp/__init__.py b/tests/ignite/metrics/nlp/__init__.py index 7cceac1455ca..e666ce106bb9 100644 --- a/tests/ignite/metrics/nlp/__init__.py +++ b/tests/ignite/metrics/nlp/__init__.py @@ -17,10 +17,8 @@ def preproc(text): self.cand_2a = preproc( "It is a guide to action which ensures that the military always obeys the commands of the party" ) - self.cand_2b = preproc("It is to insure the troops forever hearing the activity guidebook that " "party direct") - self.ref_2a = preproc( - "It is a guide to action that ensures that the military will forever heed " "Party commands" - ) + self.cand_2b = preproc("It is to insure the troops forever hearing the activity guidebook that party direct") + self.ref_2a = preproc("It is a guide to action that ensures that the military will forever heed Party commands") self.ref_2b = preproc( "It is the guiding principle which guarantees the military forces always being under the command of " "the Party" diff --git a/tests/ignite/metrics/test_accumulation.py b/tests/ignite/metrics/test_accumulation.py index d4551721ee0e..e3c1a2bf1286 100644 --- a/tests/ignite/metrics/test_accumulation.py +++ b/tests/ignite/metrics/test_accumulation.py @@ -382,14 +382,14 @@ def _test_distrib_accumulator_device(device): for metric_device in metric_devices: m = VariableAccumulation(lambda a, x: x, device=metric_device) assert m._device == metric_device - assert ( - m.accumulator.device == metric_device - ), f"{type(m.accumulator.device)}:{m.accumulator.device} vs {type(metric_device)}:{metric_device}" + assert m.accumulator.device == metric_device, ( + f"{type(m.accumulator.device)}:{m.accumulator.device} vs {type(metric_device)}:{metric_device}" + ) m.update(torch.tensor(1, device=device)) - assert ( - m.accumulator.device == metric_device - ), f"{type(m.accumulator.device)}:{m.accumulator.device} vs {type(metric_device)}:{metric_device}" + assert m.accumulator.device == metric_device, ( + f"{type(m.accumulator.device)}:{m.accumulator.device} vs {type(metric_device)}:{metric_device}" + ) def _test_apex_average(device, amp_mode, opt_level): diff --git a/tests/ignite/metrics/test_accuracy.py b/tests/ignite/metrics/test_accuracy.py index 0158739e6eac..668a40b27c8a 100644 --- a/tests/ignite/metrics/test_accuracy.py +++ b/tests/ignite/metrics/test_accuracy.py @@ -211,9 +211,9 @@ def test_multilabel_input_NHW(self): y = torch.randint(0, 2, size=(4, 5, 8, 10), device=device).long() acc.update((y_pred, y)) - assert ( - acc._num_correct.device == metric_device - ), f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + assert acc._num_correct.device == metric_device, ( + f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + ) n = acc._num_examples assert n == y.numel() / y.size(dim=1) @@ -236,9 +236,9 @@ def test_multilabel_input_NHW(self): y = torch.randint(0, 2, size=(4, 7, 10, 8), device=device).long() acc.update((y_pred, y)) - assert ( - acc._num_correct.device == metric_device - ), f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + assert acc._num_correct.device == metric_device, ( + f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + ) n = acc._num_examples assert n == y.numel() / y.size(dim=1) @@ -274,9 +274,9 @@ def test_multilabel_input_NHW(self): idx = i * batch_size acc.update((y_pred[idx : idx + batch_size], y[idx : idx + batch_size])) - assert ( - acc._num_correct.device == metric_device - ), f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + assert acc._num_correct.device == metric_device, ( + f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + ) n = acc._num_examples assert n == y.numel() / y.size(dim=1) @@ -329,9 +329,9 @@ def update(engine, i): y_true = idist.all_gather(y_true) y_preds = idist.all_gather(y_preds) - assert ( - acc._num_correct.device == metric_device - ), f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + assert acc._num_correct.device == metric_device, ( + f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + ) assert "acc" in engine.state.metrics res = engine.state.metrics["acc"] @@ -387,9 +387,9 @@ def update(engine, i): y_true = idist.all_gather(y_true) y_preds = idist.all_gather(y_preds) - assert ( - acc._num_correct.device == metric_device - ), f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + assert acc._num_correct.device == metric_device, ( + f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + ) assert "acc" in engine.state.metrics res = engine.state.metrics["acc"] @@ -409,17 +409,17 @@ def test_accumulator_device(self): for metric_device in metric_devices: acc = Accuracy(device=metric_device) assert acc._device == metric_device - assert ( - acc._num_correct.device == metric_device - ), f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + assert acc._num_correct.device == metric_device, ( + f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + ) y_pred = torch.randint(0, 2, size=(10,), device=device, dtype=torch.long) y = torch.randint(0, 2, size=(10,), device=device, dtype=torch.long) acc.update((y_pred, y)) - assert ( - acc._num_correct.device == metric_device - ), f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + assert acc._num_correct.device == metric_device, ( + f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + ) @pytest.mark.parametrize("n_epochs", [1, 2]) def test_integration_list_of_tensors_or_numbers(self, n_epochs): @@ -456,9 +456,9 @@ def update(_, i): y_true = idist.all_gather(y_true) y_preds = idist.all_gather(y_preds) - assert ( - acc._num_correct.device == metric_device - ), f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + assert acc._num_correct.device == metric_device, ( + f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + ) assert "acc" in engine.state.metrics res = engine.state.metrics["acc"] diff --git a/tests/ignite/metrics/test_classification_report.py b/tests/ignite/metrics/test_classification_report.py index 3eda0334deaa..4c069b834e4e 100644 --- a/tests/ignite/metrics/test_classification_report.py +++ b/tests/ignite/metrics/test_classification_report.py @@ -170,9 +170,9 @@ def test_metrics_result_mode(metrics_result_mode): metric = ClassificationReport(output_dict=True, metrics_result_mode=metrics_result_mode) assert isinstance(metric, MetricsLambda), "ClassificationReport should be an instance of MetricsLambda" - assert ( - metric._metrics_result_mode == metrics_result_mode - ), f"Expected metrics_result_mode to be {metrics_result_mode}" + assert metric._metrics_result_mode == metrics_result_mode, ( + f"Expected metrics_result_mode to be {metrics_result_mode}" + ) def _test_integration_multilabel(device, output_dict): diff --git a/tests/ignite/metrics/test_confusion_matrix.py b/tests/ignite/metrics/test_confusion_matrix.py index 7973ee110c59..d922480898bd 100644 --- a/tests/ignite/metrics/test_confusion_matrix.py +++ b/tests/ignite/metrics/test_confusion_matrix.py @@ -535,17 +535,17 @@ def _test_distrib_accumulator_device(device): for metric_device in metric_devices: cm = ConfusionMatrix(num_classes=3, device=metric_device) assert cm._device == metric_device - assert ( - cm.confusion_matrix.device == metric_device - ), f"{type(cm.confusion_matrix.device)}:{cm._num_correct.device} vs {type(metric_device)}:{metric_device}" + assert cm.confusion_matrix.device == metric_device, ( + f"{type(cm.confusion_matrix.device)}:{cm._num_correct.device} vs {type(metric_device)}:{metric_device}" + ) y_true, y_pred = get_y_true_y_pred() th_y_true, th_y_logits = compute_th_y_true_y_logits(y_true, y_pred) cm.update((th_y_logits, th_y_true)) - assert ( - cm.confusion_matrix.device == metric_device - ), f"{type(cm.confusion_matrix.device)}:{cm._num_correct.device} vs {type(metric_device)}:{metric_device}" + assert cm.confusion_matrix.device == metric_device, ( + f"{type(cm.confusion_matrix.device)}:{cm._num_correct.device} vs {type(metric_device)}:{metric_device}" + ) @pytest.mark.parametrize("average", [None, "samples"]) diff --git a/tests/ignite/metrics/test_loss.py b/tests/ignite/metrics/test_loss.py index 239d75197fd0..5f328d232aa5 100644 --- a/tests/ignite/metrics/test_loss.py +++ b/tests/ignite/metrics/test_loss.py @@ -184,16 +184,16 @@ def _test_distrib_accumulator_device(device, y_test_1): for metric_device in metric_devices: loss = Loss(nll_loss, device=metric_device) assert loss._device == metric_device - assert ( - loss._sum.device == metric_device - ), f"{type(loss._sum.device)}:{loss._sum.device} vs {type(metric_device)}:{metric_device}" + assert loss._sum.device == metric_device, ( + f"{type(loss._sum.device)}:{loss._sum.device} vs {type(metric_device)}:{metric_device}" + ) y_pred, y, _ = y_test_1 loss.update((y_pred, y)) - assert ( - loss._sum.device == metric_device - ), f"{type(loss._sum.device)}:{loss._sum.device} vs {type(metric_device)}:{metric_device}" + assert loss._sum.device == metric_device, ( + f"{type(loss._sum.device)}:{loss._sum.device} vs {type(metric_device)}:{metric_device}" + ) def test_sum_detached(): diff --git a/tests/ignite/metrics/test_mean_average_precision.py b/tests/ignite/metrics/test_mean_average_precision.py index 16be8b7fbb04..4a4d4ecff8eb 100644 --- a/tests/ignite/metrics/test_mean_average_precision.py +++ b/tests/ignite/metrics/test_mean_average_precision.py @@ -174,7 +174,6 @@ def test_distrib_integration(distributed, data_type, n_epochs): metric_devices.append(device) for metric_device in metric_devices: - y_true_size = ( (n_iters * batch_size, 3, 2) if data_type != "multilabel" else (n_iters * batch_size, n_classes, 3, 2) ) diff --git a/tests/ignite/metrics/test_multilabel_confusion_matrix.py b/tests/ignite/metrics/test_multilabel_confusion_matrix.py index a67284501467..c2e8bd049e39 100644 --- a/tests/ignite/metrics/test_multilabel_confusion_matrix.py +++ b/tests/ignite/metrics/test_multilabel_confusion_matrix.py @@ -198,16 +198,16 @@ def _test_distrib_accumulator_device(device): for metric_device in metric_devices: cm = MultiLabelConfusionMatrix(num_classes=3, device=metric_device) assert cm._device == metric_device - assert ( - cm.confusion_matrix.device == metric_device - ), f"{type(cm.confusion_matrix.device)}:{cm._num_correct.device} vs {type(metric_device)}:{metric_device}" + assert cm.confusion_matrix.device == metric_device, ( + f"{type(cm.confusion_matrix.device)}:{cm._num_correct.device} vs {type(metric_device)}:{metric_device}" + ) y_true, y_pred = get_y_true_y_pred() cm.update((torch.tensor(y_pred), torch.tensor(y_true))) - assert ( - cm.confusion_matrix.device == metric_device - ), f"{type(cm.confusion_matrix.device)}:{cm._num_correct.device} vs {type(metric_device)}:{metric_device}" + assert cm.confusion_matrix.device == metric_device, ( + f"{type(cm.confusion_matrix.device)}:{cm._num_correct.device} vs {type(metric_device)}:{metric_device}" + ) def test_simple_2D_input(available_device): diff --git a/tests/ignite/metrics/test_precision.py b/tests/ignite/metrics/test_precision.py index 06b8f6eafd5c..7bd80540f5de 100644 --- a/tests/ignite/metrics/test_precision.py +++ b/tests/ignite/metrics/test_precision.py @@ -457,15 +457,15 @@ def test_accumulator_device(self, average): assert pr._updated is True - assert ( - pr._numerator.device == metric_device - ), f"{type(pr._numerator.device)}:{pr._numerator.device} vs {type(metric_device)}:{metric_device}" + assert pr._numerator.device == metric_device, ( + f"{type(pr._numerator.device)}:{pr._numerator.device} vs {type(metric_device)}:{metric_device}" + ) if average != "samples": # For average='samples', `_denominator` is of type `int` so it has not `device` member. - assert ( - pr._denominator.device == metric_device - ), f"{type(pr._denominator.device)}:{pr._denominator.device} vs {type(metric_device)}:{metric_device}" + assert pr._denominator.device == metric_device, ( + f"{type(pr._denominator.device)}:{pr._denominator.device} vs {type(metric_device)}:{metric_device}" + ) if average == "weighted": assert pr._weight.device == metric_device, f"{type(pr._weight.device)}:{pr._weight.device} vs " @@ -490,15 +490,15 @@ def test_multilabel_accumulator_device(self, average): assert pr._updated is True - assert ( - pr._numerator.device == metric_device - ), f"{type(pr._numerator.device)}:{pr._numerator.device} vs {type(metric_device)}:{metric_device}" + assert pr._numerator.device == metric_device, ( + f"{type(pr._numerator.device)}:{pr._numerator.device} vs {type(metric_device)}:{metric_device}" + ) if average != "samples": # For average='samples', `_denominator` is of type `int` so it has not `device` member. - assert ( - pr._denominator.device == metric_device - ), f"{type(pr._denominator.device)}:{pr._denominator.device} vs {type(metric_device)}:{metric_device}" + assert pr._denominator.device == metric_device, ( + f"{type(pr._denominator.device)}:{pr._denominator.device} vs {type(metric_device)}:{metric_device}" + ) if average == "weighted": assert pr._weight.device == metric_device, f"{type(pr._weight.device)}:{pr._weight.device} vs " diff --git a/tests/ignite/metrics/test_recall.py b/tests/ignite/metrics/test_recall.py index fc7f29236fa8..90e78d1a82ad 100644 --- a/tests/ignite/metrics/test_recall.py +++ b/tests/ignite/metrics/test_recall.py @@ -263,7 +263,6 @@ def to_numpy_multilabel(y): @pytest.mark.parametrize("n_times", range(3)) @pytest.mark.parametrize("average", [None, False, "macro", "micro", "weighted", "samples"]) def test_multilabel_input(n_times, available_device, average, test_data_multilabel): - re = Recall(average=average, is_multilabel=True, device=available_device) assert re._device == torch.device(available_device) assert re._updated is False @@ -459,15 +458,15 @@ def test_accumulator_device(self, average): assert re._updated is True - assert ( - re._numerator.device == metric_device - ), f"{type(re._numerator.device)}:{re._numerator.device} vs {type(metric_device)}:{metric_device}" + assert re._numerator.device == metric_device, ( + f"{type(re._numerator.device)}:{re._numerator.device} vs {type(metric_device)}:{metric_device}" + ) if average != "samples": # For average='samples', `_denominator` is of type `int` so it has not `device` member. - assert ( - re._denominator.device == metric_device - ), f"{type(re._denominator.device)}:{re._denominator.device} vs {type(metric_device)}:{metric_device}" + assert re._denominator.device == metric_device, ( + f"{type(re._denominator.device)}:{re._denominator.device} vs {type(metric_device)}:{metric_device}" + ) if average == "weighted": assert re._weight.device == metric_device, f"{type(re._weight.device)}:{re._weight.device} vs " @@ -493,15 +492,15 @@ def test_multilabel_accumulator_device(self, average): assert re._updated is True - assert ( - re._numerator.device == metric_device - ), f"{type(re._numerator.device)}:{re._numerator.device} vs {type(metric_device)}:{metric_device}" + assert re._numerator.device == metric_device, ( + f"{type(re._numerator.device)}:{re._numerator.device} vs {type(metric_device)}:{metric_device}" + ) if average != "samples": # For average='samples', `_denominator` is of type `int` so it has not `device` member. - assert ( - re._denominator.device == metric_device - ), f"{type(re._denominator.device)}:{re._denominator.device} vs {type(metric_device)}:{metric_device}" + assert re._denominator.device == metric_device, ( + f"{type(re._denominator.device)}:{re._denominator.device} vs {type(metric_device)}:{metric_device}" + ) if average == "weighted": assert re._weight.device == metric_device, f"{type(re._weight.device)}:{re._weight.device} vs " diff --git a/tests/ignite/metrics/test_running_average.py b/tests/ignite/metrics/test_running_average.py index dc434bda636d..523b11b3dd86 100644 --- a/tests/ignite/metrics/test_running_average.py +++ b/tests/ignite/metrics/test_running_average.py @@ -293,9 +293,9 @@ def running_avg_output_update(engine): @trainer.on(usage.COMPLETED) def assert_equal_running_avg_output_values(engine): it = engine.state.iteration - assert ( - engine.state.running_avg_output == engine.state.metrics["running_avg_output"] - ), f"{it}: {engine.state.running_avg_output} vs {engine.state.metrics['running_avg_output']}" + assert engine.state.running_avg_output == engine.state.metrics["running_avg_output"], ( + f"{it}: {engine.state.running_avg_output} vs {engine.state.metrics['running_avg_output']}" + ) trainer.run(data, max_epochs=3) @@ -365,9 +365,9 @@ def assert_equal_running_avg_acc_values(engine): if not isinstance(usage, RunningEpochWise) or ( (engine.state.iteration > 1) and ((engine.state.iteration % n_iters) == 1) ): - assert ( - engine.state.running_avg_acc == engine.state.metrics["running_avg_accuracy"] - ), f"{engine.state.running_avg_acc} vs {engine.state.metrics['running_avg_accuracy']}" + assert engine.state.running_avg_acc == engine.state.metrics["running_avg_accuracy"], ( + f"{engine.state.running_avg_acc} vs {engine.state.metrics['running_avg_accuracy']}" + ) trainer.run(data, max_epochs=3) @@ -391,6 +391,6 @@ def test_accumulator_device(self): avg.update(torch.tensor(1.0, device=device)) avg.compute() - assert ( - avg._value.device == metric_device - ), f"{type(avg._value.device)}:{avg._value.device} vs {type(metric_device)}:{metric_device}" + assert avg._value.device == metric_device, ( + f"{type(avg._value.device)}:{avg._value.device} vs {type(metric_device)}:{metric_device}" + ) diff --git a/tests/ignite/metrics/test_top_k_categorical_accuracy.py b/tests/ignite/metrics/test_top_k_categorical_accuracy.py index c05fd1451b81..4c8dea00570c 100644 --- a/tests/ignite/metrics/test_top_k_categorical_accuracy.py +++ b/tests/ignite/metrics/test_top_k_categorical_accuracy.py @@ -112,17 +112,17 @@ def _test_distrib_accumulator_device(device): for metric_device in metric_devices: acc = TopKCategoricalAccuracy(2, device=metric_device) assert acc._device == metric_device - assert ( - acc._num_correct.device == metric_device - ), f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + assert acc._num_correct.device == metric_device, ( + f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + ) y_pred = torch.tensor([[0.2, 0.4, 0.6, 0.8], [0.8, 0.6, 0.4, 0.2]]) y = torch.ones(2).long() acc.update((y_pred, y)) - assert ( - acc._num_correct.device == metric_device - ), f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + assert acc._num_correct.device == metric_device, ( + f"{type(acc._num_correct.device)}:{acc._num_correct.device} vs {type(metric_device)}:{metric_device}" + ) @pytest.mark.distributed From 325a20f113bed3fa320b816ec9e2e18f6dd3b827 Mon Sep 17 00:00:00 2001 From: blanky Date: Thu, 26 Mar 2026 01:40:29 +0530 Subject: [PATCH 3/4] fix: correct assertion syntax in test_output_handler_metric_names --- tests/ignite/handlers/test_neptune_logger.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/tests/ignite/handlers/test_neptune_logger.py b/tests/ignite/handlers/test_neptune_logger.py index 67b1e058a5a5..2abea6c516f2 100644 --- a/tests/ignite/handlers/test_neptune_logger.py +++ b/tests/ignite/handlers/test_neptune_logger.py @@ -200,13 +200,13 @@ def test_output_handler_metric_names(): wrapper(mock_engine, logger, Events.ITERATION_STARTED) - (assert_logger_called_once_with(logger, "tag/a", 123),) - (assert_logger_called_once_with(logger, "tag/b/c/0", 2.34),) - (assert_logger_called_once_with(logger, "tag/b/c/1/d", 1),) - (assert_logger_called_once_with(logger, "tag/c/0", 22),) - (assert_logger_called_once_with(logger, "tag/c/1/0", 33),) - (assert_logger_called_once_with(logger, "tag/c/1/1", -5.5),) - (assert_logger_called_once_with(logger, "tag/c/2/e", 32.1),) + (assert_logger_called_once_with(logger, "tag/a", 123)) + (assert_logger_called_once_with(logger, "tag/b/c/0", 2.34)) + (assert_logger_called_once_with(logger, "tag/b/c/1/d", 1)) + (assert_logger_called_once_with(logger, "tag/c/0", 22)) + (assert_logger_called_once_with(logger, "tag/c/1/0", 33)) + (assert_logger_called_once_with(logger, "tag/c/1/1", -5.5)) + (assert_logger_called_once_with(logger, "tag/c/2/e", 32.1)) logger.stop() From f041dd46de0d568a42f362b9e570de41a806ec26 Mon Sep 17 00:00:00 2001 From: vfdev Date: Wed, 25 Mar 2026 23:20:35 +0100 Subject: [PATCH 4/4] Apply suggestion from @vfdev-5 --- tests/ignite/handlers/test_neptune_logger.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/tests/ignite/handlers/test_neptune_logger.py b/tests/ignite/handlers/test_neptune_logger.py index 2abea6c516f2..06334856e1ab 100644 --- a/tests/ignite/handlers/test_neptune_logger.py +++ b/tests/ignite/handlers/test_neptune_logger.py @@ -200,13 +200,13 @@ def test_output_handler_metric_names(): wrapper(mock_engine, logger, Events.ITERATION_STARTED) - (assert_logger_called_once_with(logger, "tag/a", 123)) - (assert_logger_called_once_with(logger, "tag/b/c/0", 2.34)) - (assert_logger_called_once_with(logger, "tag/b/c/1/d", 1)) - (assert_logger_called_once_with(logger, "tag/c/0", 22)) - (assert_logger_called_once_with(logger, "tag/c/1/0", 33)) - (assert_logger_called_once_with(logger, "tag/c/1/1", -5.5)) - (assert_logger_called_once_with(logger, "tag/c/2/e", 32.1)) + assert_logger_called_once_with(logger, "tag/a", 123) + assert_logger_called_once_with(logger, "tag/b/c/0", 2.34) + assert_logger_called_once_with(logger, "tag/b/c/1/d", 1) + assert_logger_called_once_with(logger, "tag/c/0", 22) + assert_logger_called_once_with(logger, "tag/c/1/0", 33) + assert_logger_called_once_with(logger, "tag/c/1/1", -5.5) + assert_logger_called_once_with(logger, "tag/c/2/e", 32.1) logger.stop()