From 79ad249f1519b1440ff8d26df430d010a292b4ff Mon Sep 17 00:00:00 2001 From: tnw-chatc Date: Tue, 5 Aug 2025 07:41:33 -0700 Subject: [PATCH 1/2] Implement arbitrary input channels SimCLR --- src/hyrax/hyrax_default_config.toml | 3 +++ src/hyrax/models/simclr.py | 22 +++++++++++++++++++--- 2 files changed, 22 insertions(+), 3 deletions(-) diff --git a/src/hyrax/hyrax_default_config.toml b/src/hyrax/hyrax_default_config.toml index b69058a6b..1ad08324e 100644 --- a/src/hyrax/hyrax_default_config.toml +++ b/src/hyrax/hyrax_default_config.toml @@ -101,6 +101,9 @@ projection_dimension = 128 # The scalar temperature parameter for its loss function, NTXentLoss, for SimCLR temperature = 0.5 +# The number of the input channels. +input_channels = 3 + # The probability of applying horizontal flip augmentation for SimCLR horizontal_flip_probability = 0.5 diff --git a/src/hyrax/models/simclr.py b/src/hyrax/models/simclr.py index fbe2081b3..cc7822b3e 100644 --- a/src/hyrax/models/simclr.py +++ b/src/hyrax/models/simclr.py @@ -13,7 +13,7 @@ class NTXentLoss(nn.Module): """Normalized Temperature-scaled Cross Entropy Loss. Based on Chen, 2020""" - def __init__(self, temperature=0.1): + def __init__(self, temperature): super().__init__() self.temperature = temperature self.criterion = nn.CrossEntropyLoss(reduction="sum") @@ -22,6 +22,7 @@ def forward(self, z_i, z_j): """Forward function of NTXentLoss. Based on Chen, 2020. Loss is calculated from representations from two augmented views of the same batch. """ + batch_size = z_i.shape[0] device = z_i.device @@ -72,8 +73,22 @@ def __init__(self, config, shape): self.shape = shape proj_dim = config["model"]["SimCLR"]["projection_dimension"] temperature = config["model"]["SimCLR"]["temperature"] + input_channels = config["model"]["SimCLR"]["input_channels"] backbone = models.resnet18(pretrained=False) + + # Create a new conv layer with same out_channels, kernel size, stride, padding, etc + # But with the new in_channels + old_conv = backbone.conv1 + backbone.conv1 = nn.Conv2d( + in_channels=input_channels, + out_channels=old_conv.out_channels, + kernel_size=old_conv.kernel_size, + stride=old_conv.stride, + padding=old_conv.padding, + bias=old_conv.bias is not None, + ) + backbone.fc = nn.Identity() self.backbone = backbone @@ -82,7 +97,8 @@ def __init__(self, config, shape): nn.ReLU(inplace=True), nn.Linear(512, proj_dim), ) - self.criterion = NTXentLoss(temperature) + # TODO: Make sure to revisit this and properly implement custom criterion + self.the_criterion = NTXentLoss(temperature) def forward(self, x): feats = self.backbone(x) @@ -111,7 +127,7 @@ def train_step(self, x): z1 = self.forward(x1) z2 = self.forward(x2) - loss = self.criterion(z1, z2) + loss = self.the_criterion(z1, z2) self.optimizer.zero_grad() loss.backward() self.optimizer.step() From 0145ecb2a1c6ae57860f4f9ea20d59e6529c4229 Mon Sep 17 00:00:00 2001 From: tnw-chatc Date: Sun, 10 Aug 2025 23:22:49 -0700 Subject: [PATCH 2/2] Add SimCLRv2 for arbitrary dimensional input --- src/hyrax/hyrax_default_config.toml | 9 ++++ src/hyrax/models/__init__.py | 3 +- src/hyrax/models/simclr.py | 70 ++++++++++++++++++++++++++--- 3 files changed, 75 insertions(+), 7 deletions(-) diff --git a/src/hyrax/hyrax_default_config.toml b/src/hyrax/hyrax_default_config.toml index 1ad08324e..5898697f5 100644 --- a/src/hyrax/hyrax_default_config.toml +++ b/src/hyrax/hyrax_default_config.toml @@ -123,6 +123,15 @@ gaussian_blur_kernel_size = 9 # The sigma range used in Gaussian blur augmentation for SimCLR gaussian_blur_sigma_range = [0.1, 2.0] +# The maximum rotation angle (in degrees) augmentation for SimCLR +rotation_range = 180 + +# The mean of the distribution for Gaussian noise augmentation for SimCLR +gaussian_noise_mean = 0 + +# The standard deviation of the distribution for Gaussian noise augmentation for SimCLR +gaussian_noise_sigma = 0.2 + [criterion] # The name of the built-in criterion to use or the import path to an external criterion diff --git a/src/hyrax/models/__init__.py b/src/hyrax/models/__init__.py index 1e7dc5ad4..570e8945f 100644 --- a/src/hyrax/models/__init__.py +++ b/src/hyrax/models/__init__.py @@ -9,7 +9,7 @@ from .hyrax_cnn import HyraxCNN from .hyrax_loopback import HyraxLoopback from .model_registry import hyrax_model -from .simclr import SimCLR +from .simclr import SimCLR, SimCLRv2 __all__ = [ "hyrax_model", @@ -21,4 +21,5 @@ "HSCDCAE", "ImageDCAE", "SimCLR", + "SimCLRv2", ] diff --git a/src/hyrax/models/simclr.py b/src/hyrax/models/simclr.py index cc7822b3e..48d281944 100644 --- a/src/hyrax/models/simclr.py +++ b/src/hyrax/models/simclr.py @@ -5,7 +5,7 @@ import torch.nn as nn import torch.nn.functional as F # noqa N812 import torchvision.models as models -import torchvision.transforms as T # noqa N812 +import torchvision.transforms.v2 as T # noqa N812 from hyrax.models.model_registry import hyrax_model @@ -67,6 +67,64 @@ def __call__(self, x): class SimCLR(nn.Module): """SimCLR model. Implementation based on Chen, 2020""" + def __init__(self, config, shape): + super().__init__() + self.config = config + self.shape = shape + proj_dim = config["model"]["SimCLR"]["projection_dimension"] + temperature = config["model"]["SimCLR"]["temperature"] + + backbone = models.resnet18(pretrained=False) + backbone.fc = nn.Identity() + self.backbone = backbone + + self.projection_head = nn.Sequential( + nn.Linear(512, 512), + nn.ReLU(inplace=True), + nn.Linear(512, proj_dim), + ) + # TODO: Make sure to revisit this and properly implement custom criterion + self.the_criterion = NTXentLoss(temperature) + + def forward(self, x): + feats = self.backbone(x) + return self.projection_head(feats) + + def train_step(self, x): + aug = T.Compose( + [ + T.RandomResizedCrop(size=x.shape[-1]), + T.RandomHorizontalFlip(self.config["model"]["SimCLR"]["horizontal_flip_probability"]), + T.RandomApply( + [PositiveRescale(T.ColorJitter(*self.config["model"]["SimCLR"]["color_jitter_params"]))], + p=self.config["model"]["SimCLR"]["color_jitter_probability"], + ), + T.RandomGrayscale(p=self.config["model"]["SimCLR"]["grayscale_probability"]), + T.GaussianBlur( + kernel_size=self.config["model"]["SimCLR"]["gaussian_blur_kernel_size"], + sigma=self.config["model"]["SimCLR"]["gaussian_blur_sigma_range"], + ), + ] + ) + + x1 = torch.stack([aug(img) for img in x]) + x2 = torch.stack([aug(img) for img in x]) + + z1 = self.forward(x1) + z2 = self.forward(x2) + + loss = self.the_criterion(z1, z2) + self.optimizer.zero_grad() + loss.backward() + self.optimizer.step() + return {"loss": loss.item()} + + +@hyrax_model +class SimCLRv2(nn.Module): + """Modified SimCLR model with new augmentation routine and compatible + with an arbitrary number of input channels""" + def __init__(self, config, shape): super().__init__() self.config = config @@ -109,15 +167,15 @@ def train_step(self, x): [ T.RandomResizedCrop(size=x.shape[-1]), T.RandomHorizontalFlip(self.config["model"]["SimCLR"]["horizontal_flip_probability"]), - T.RandomApply( - [PositiveRescale(T.ColorJitter(*self.config["model"]["SimCLR"]["color_jitter_params"]))], - p=self.config["model"]["SimCLR"]["color_jitter_probability"], - ), - T.RandomGrayscale(p=self.config["model"]["SimCLR"]["grayscale_probability"]), + T.RandomRotation(self.config["model"]["SimCLR"]["rotation_range"]), T.GaussianBlur( kernel_size=self.config["model"]["SimCLR"]["gaussian_blur_kernel_size"], sigma=self.config["model"]["SimCLR"]["gaussian_blur_sigma_range"], ), + T.GaussianNoise( + mean=self.config["model"]["SimCLR"]["gaussian_noise_mean"], + sigma=self.config["model"]["SimCLR"]["gaussian_noise_sigma"], + ), ] )