From a21f2c1cce780195d1f4c64e6c10fe76694caa26 Mon Sep 17 00:00:00 2001 From: kir_7 Date: Sun, 5 Jul 2026 22:30:01 +0530 Subject: [PATCH 1/7] added stop and included in ijepa --- lightly/models/modules/ijepa.py | 18 ++++++++ lightly/models/modules/ijepa_timm.py | 12 ++++++ lightly/models/utils.py | 43 +++++++++++++++++++ .../test_stochastic_positional_embedding.py | 31 +++++++++++++ 4 files changed, 104 insertions(+) create mode 100644 tests/utils/test_stochastic_positional_embedding.py diff --git a/lightly/models/modules/ijepa.py b/lightly/models/modules/ijepa.py index 7889dadd1..f112efa3d 100644 --- a/lightly/models/modules/ijepa.py +++ b/lightly/models/modules/ijepa.py @@ -79,6 +79,11 @@ def __init__( torch.from_numpy(predictor_pos_embed).float().unsqueeze(0) ) + self.use_stop = kwargs.get( + "use_stop", False + ) # pass use stop embeddings as additional args, default to False + self.noise_std = kwargs.get("noise_std", 0.25) # default 0.25 + @classmethod def from_vit_encoder(cls, vit_encoder, num_patches): """Creates an I-JEPA predictor backbone (multi-head attention and layernorm) from a torchvision ViT encoder. @@ -134,6 +139,7 @@ def forward(self, x, masks_x, masks): if not isinstance(masks, list): masks = [masks] + noise_dim = x.shape[-1] B = len(x) // len(masks_x) x = self.predictor_embed(x) x_pos_embed = self.predictor_pos_embed.repeat(B, 1, 1) @@ -144,9 +150,21 @@ def forward(self, x, masks_x, masks): pos_embs = self.predictor_pos_embed.repeat(B, 1, 1) pos_embs = utils.apply_masks(pos_embs, masks) pos_embs = utils.repeat_interleave_batch(pos_embs, B, repeat=len(masks_x)) + + # we add the stochastic positional embedding here: + # use self.predictor_embed as the projector + pos_embs = utils.add_stochastic_positional_noise( + pos_embs, + self.predictor_embed, + noise_dim, + noise_std=self.noise_std, + enabled=self.use_stop, + ) + pred_tokens = self.mask_token.repeat(pos_embs.size(0), pos_embs.size(1), 1) pred_tokens += pos_embs + x = x.repeat(len(masks), 1, 1) x = torch.cat([x, pred_tokens], dim=1) diff --git a/lightly/models/modules/ijepa_timm.py b/lightly/models/modules/ijepa_timm.py index ebba789ab..4b8fcb3d7 100644 --- a/lightly/models/modules/ijepa_timm.py +++ b/lightly/models/modules/ijepa_timm.py @@ -123,6 +123,7 @@ def forward( len_masks_x = len(masks_x) if isinstance(masks_x, list) else 1 len_masks = len(masks) if isinstance(masks, list) else 1 + noise_dim = x.shape[-1] B = len(x) // len_masks_x x = self.predictor_embed(x) x_pos_embed = self.predictor_pos_embed.repeat(B, 1, 1) @@ -136,6 +137,17 @@ def forward( pred_tokens = self.mask_token.repeat(pos_embs.size(0), pos_embs.size(1), 1) pred_tokens += pos_embs + + # we add the stochastic positional embedding here: + # use self.predictor_embed as the projector + pred_tokens = utils.add_stochastic_positional_noise( + pred_tokens, + self.predictor_embed, + noise_dim, + noise_std=self.noise_std, + enabled=self.use_stop, + ) + x = x.repeat(len_masks, 1, 1) x = torch.cat([x, pred_tokens], dim=1) diff --git a/lightly/models/utils.py b/lightly/models/utils.py index eb6c17533..31d3069f3 100644 --- a/lightly/models/utils.py +++ b/lightly/models/utils.py @@ -1316,3 +1316,46 @@ def apply_masks(x: Tensor, masks: Tensor | list[Tensor]) -> Tensor: mask_keep = m.unsqueeze(-1).repeat(1, 1, x.size(-1)) all_x += [torch.gather(x, dim=1, index=mask_keep)] return torch.cat(all_x, dim=0) + + +def add_stochastic_positional_noise( + pos_embeddings: Tensor, + projection: Module, + noise_dim: int, + noise_std: float = 0.25, + enabled: bool = False, +) -> Tensor: + """Adds stochastic noise to positional embeddings. + + [0]. https://arxiv.org/pdf/2308.00566 + [1]. https://github.com/amirbar/StoP/blob/main/src/deit.py + + Args: + pos_embeddings: + Positional embeddings of shape + ``(batch_size, num_tokens, predictor_embed_dim)``. + projection: + Matrix A used to project gaussian noise to the pos_embedding + dimension. + noise_dim: + Dimension of the sampled gaussian noise before projection. + noise_std: + Standard deviation of the gaussian noise. + enabled: + If False, returns ``pos_embeddings`` unchanged. + + Returns: + Positional embeddings with optional gaussian noise added. + """ + if not enabled or noise_std == 0.0: + return pos_embeddings + + noise = torch.normal( + mean=0.0, + std=noise_std, + size=(pos_embeddings.shape[0], pos_embeddings.shape[1], noise_dim), + device=pos_embeddings.device, + dtype=pos_embeddings.dtype, + ) + + return pos_embeddings + projection(noise) diff --git a/tests/utils/test_stochastic_positional_embedding.py b/tests/utils/test_stochastic_positional_embedding.py new file mode 100644 index 000000000..19bfe0958 --- /dev/null +++ b/tests/utils/test_stochastic_positional_embedding.py @@ -0,0 +1,31 @@ +import torch + +from lightly.models import utils + + +def test_add_stochastic_positional_noise_disabled() -> None: + projection = torch.nn.Linear(8, 4) + pos_embeddings = torch.randn(2, 3, 4) + + out = utils.add_stochastic_positional_noise( + pos_embeddings=pos_embeddings, + projection=projection, + noise_dim=8, + enabled=False, + ) + + assert torch.equal(out, pos_embeddings) + + +def test_add_stochastic_positional_noise_enabled_shape() -> None: + projection = torch.nn.Linear(8, 4) + pos_embeddings = torch.randn(2, 3, 4) + + out = utils.add_stochastic_positional_noise( + pos_embeddings=pos_embeddings, + projection=projection, + noise_dim=8, + enabled=True, + ) + + assert out.shape == pos_embeddings.shape From 9c27f5f63ccd3ffa8134c8644747509eba2eeff7 Mon Sep 17 00:00:00 2001 From: kir_7 Date: Mon, 6 Jul 2026 23:40:56 +0530 Subject: [PATCH 2/7] fixed failing tests for ijepa_timm; extended ijepa_timm to test stop --- lightly/models/modules/ijepa_timm.py | 5 ++++ tests/models/modules/test_ijepa_timm.py | 39 ++++++++++++++++--------- 2 files changed, 30 insertions(+), 14 deletions(-) diff --git a/lightly/models/modules/ijepa_timm.py b/lightly/models/modules/ijepa_timm.py index 4b8fcb3d7..adfe57d16 100644 --- a/lightly/models/modules/ijepa_timm.py +++ b/lightly/models/modules/ijepa_timm.py @@ -61,6 +61,8 @@ def __init__( proj_drop_rate: float = 0.0, attn_drop_rate: float = 0.0, norm_layer: Callable[..., nn.Module] = partial(nn.LayerNorm, eps=1e-6), + use_stop: bool = False, + noise_std: float = 0.25, ): """Initializes the IJEPAPredictorTIMM with the specified dimensions.""" super().__init__() @@ -97,6 +99,9 @@ def __init__( ] ) + self.use_stop = use_stop + self.noise_std = noise_std + def forward( self, x: Tensor, diff --git a/tests/models/modules/test_ijepa_timm.py b/tests/models/modules/test_ijepa_timm.py index f41a489cf..6cc1b9ccf 100644 --- a/tests/models/modules/test_ijepa_timm.py +++ b/tests/models/modules/test_ijepa_timm.py @@ -14,8 +14,10 @@ from lightly.models.modules import IJEPAPredictorTIMM -class TestIJEPAPredictorTIMM(unittest.TestCase): - def test_init(self) -> None: +class TestIJEPAPredictorTIMM: + @pytest.mark.parametrize("use_stop", [True, False]) + @pytest.mark.parametrize("noise_std", [0.0, 0.1]) + def test_init(self, use_stop: bool, noise_std: float) -> None: IJEPAPredictorTIMM( num_patches=196, depth=2, @@ -26,10 +28,17 @@ def test_init(self) -> None: mlp_ratio=4.0, proj_drop_rate=0.0, attn_drop_rate=0.0, + use_stop=use_stop, + noise_std=noise_std, ) def _test_forward( - self, device: torch.device, batch_size: int = 4, seed: int = 0 + self, + device: torch.device, + use_stop: bool, + noise_std: float, + batch_size: int = 4, + seed: int = 0, ) -> None: torch.manual_seed(seed) num_patches = 196 # 14x14 patches @@ -48,6 +57,8 @@ def _test_forward( mlp_ratio=4.0, proj_drop_rate=0.0, attn_drop_rate=0.0, + use_stop=use_stop, + noise_std=noise_std, ).to(device) x = torch.randn(batch_size, num_patches, mlp_dim, device=device) @@ -56,16 +67,16 @@ def _test_forward( predictions = predictor(x, masks_x, masks) - # output shape must be correct - expected_shape = [batch_size, num_patches, mlp_dim] - self.assertListEqual(list(predictions.shape), expected_shape) + assert list(predictions.shape) == [batch_size, num_patches, mlp_dim] + assert torch.all(torch.isfinite(predictions)) - # output must have reasonable numbers - self.assertTrue(torch.all(torch.isfinite(predictions))) + @pytest.mark.parametrize("use_stop", [True, False]) + @pytest.mark.parametrize("noise_std", [0.0, 0.1]) + def test_forward(self, use_stop: bool, noise_std: float) -> None: + self._test_forward(torch.device("cpu"), use_stop, noise_std) - def test_forward(self) -> None: - self._test_forward(torch.device("cpu")) - - @unittest.skipUnless(torch.cuda.is_available(), "CUDA not available.") - def test_forward_cuda(self) -> None: - self._test_forward(torch.device("cuda")) + @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available.") + @pytest.mark.parametrize("use_stop", [True, False]) + @pytest.mark.parametrize("noise_std", [0.0, 0.1]) + def test_forward_cuda(self, use_stop: bool, noise_std: float) -> None: + self._test_forward(torch.device("cuda"), use_stop, noise_std) From 2967d4f42fb497ca09fd7e8fc64341ea6712dda0 Mon Sep 17 00:00:00 2001 From: kir_7 Date: Thu, 9 Jul 2026 23:05:38 +0530 Subject: [PATCH 3/7] StoP: removed use_stop, fixed docstrings --- lightly/models/modules/ijepa.py | 12 ++++---- lightly/models/modules/ijepa_timm.py | 19 +++++++------ lightly/models/utils.py | 28 ++++++++----------- tests/models/modules/test_ijepa_timm.py | 19 ++++--------- .../test_stochastic_positional_embedding.py | 10 +++---- 5 files changed, 36 insertions(+), 52 deletions(-) diff --git a/lightly/models/modules/ijepa.py b/lightly/models/modules/ijepa.py index f112efa3d..2cd6f11b9 100644 --- a/lightly/models/modules/ijepa.py +++ b/lightly/models/modules/ijepa.py @@ -39,6 +39,9 @@ class IJEPAPredictor(vision_transformer.Encoder): Percentage of elements set to zero after the MLP in the transformer. attention_dropout: Percentage of elements set to zero after the attention head. + noise_std: + Standard deviation of the Gaussian noise added to positional embeddings. + Default ``0.0`` to disable stochastic positional embeddings. """ def __init__( @@ -53,6 +56,7 @@ def __init__( dropout: float, attention_dropout: float, norm_layer: Callable[..., torch.nn.Module] = partial(nn.LayerNorm, eps=1e-6), + noise_std: float = 0.0, **kwargs, ): """Initializes the IJEPAPredictor with the specified dimensions.""" @@ -79,10 +83,7 @@ def __init__( torch.from_numpy(predictor_pos_embed).float().unsqueeze(0) ) - self.use_stop = kwargs.get( - "use_stop", False - ) # pass use stop embeddings as additional args, default to False - self.noise_std = kwargs.get("noise_std", 0.25) # default 0.25 + self.noise_std = noise_std @classmethod def from_vit_encoder(cls, vit_encoder, num_patches): @@ -158,7 +159,6 @@ def forward(self, x, masks_x, masks): self.predictor_embed, noise_dim, noise_std=self.noise_std, - enabled=self.use_stop, ) pred_tokens = self.mask_token.repeat(pos_embs.size(0), pos_embs.size(1), 1) @@ -184,7 +184,7 @@ class IJEPAEncoder(vision_transformer.Encoder): Encodes patch embeddings. Code inspired by [1]. - - [0]: Joint-Embedding Predictive Architecture, 2023, https://arxiv.org/abs/2301.08243 + - [`0]: Joint-Embedding Predictive Architecture, 2023, https://arxiv.org/abs/2301.08243 - [1]: https://github.com/facebookresearch/ijepa Attributes: diff --git a/lightly/models/modules/ijepa_timm.py b/lightly/models/modules/ijepa_timm.py index adfe57d16..9bd840fef 100644 --- a/lightly/models/modules/ijepa_timm.py +++ b/lightly/models/modules/ijepa_timm.py @@ -46,6 +46,9 @@ class IJEPAPredictorTIMM(nn.Module): # type: ignore[misc] Percentage of elements set to zero after the attention head. norm_layer: Normalization layer. + noise_std: + Standard deviation of the Gaussian noise added to positional embeddings. + Default ``0.0`` to disable stochastic positional embeddings. """ def __init__( @@ -61,8 +64,7 @@ def __init__( proj_drop_rate: float = 0.0, attn_drop_rate: float = 0.0, norm_layer: Callable[..., nn.Module] = partial(nn.LayerNorm, eps=1e-6), - use_stop: bool = False, - noise_std: float = 0.25, + noise_std: float = 0.0, ): """Initializes the IJEPAPredictorTIMM with the specified dimensions.""" super().__init__() @@ -99,7 +101,6 @@ def __init__( ] ) - self.use_stop = use_stop self.noise_std = noise_std def forward( @@ -139,20 +140,20 @@ def forward( pos_embs = self.predictor_pos_embed.repeat(B, 1, 1) pos_embs = utils.apply_masks(pos_embs, masks) pos_embs = utils.repeat_interleave_batch(pos_embs, B, repeat=len_masks_x) - pred_tokens = self.mask_token.repeat(pos_embs.size(0), pos_embs.size(1), 1) - - pred_tokens += pos_embs # we add the stochastic positional embedding here: # use self.predictor_embed as the projector - pred_tokens = utils.add_stochastic_positional_noise( - pred_tokens, + pos_embs = utils.add_stochastic_positional_noise( + pos_embs, self.predictor_embed, noise_dim, noise_std=self.noise_std, - enabled=self.use_stop, ) + pred_tokens = self.mask_token.repeat(pos_embs.size(0), pos_embs.size(1), 1) + + pred_tokens += pos_embs + x = x.repeat(len_masks, 1, 1) x = torch.cat([x, pred_tokens], dim=1) diff --git a/lightly/models/utils.py b/lightly/models/utils.py index 31d3069f3..a562e0e6f 100644 --- a/lightly/models/utils.py +++ b/lightly/models/utils.py @@ -1322,32 +1322,26 @@ def add_stochastic_positional_noise( pos_embeddings: Tensor, projection: Module, noise_dim: int, - noise_std: float = 0.25, - enabled: bool = False, + noise_std: float = 0.0, ) -> Tensor: """Adds stochastic noise to positional embeddings. - [0]. https://arxiv.org/pdf/2308.00566 - [1]. https://github.com/amirbar/StoP/blob/main/src/deit.py + - [0]: https://arxiv.org/pdf/2308.00566 + - [1]: https://github.com/amirbar/StoP/blob/main/src/deit.py Args: - pos_embeddings: - Positional embeddings of shape + pos_embeddings: Positional embeddings of shape ``(batch_size, num_tokens, predictor_embed_dim)``. - projection: - Matrix A used to project gaussian noise to the pos_embedding - dimension. - noise_dim: - Dimension of the sampled gaussian noise before projection. - noise_std: - Standard deviation of the gaussian noise. - enabled: - If False, returns ``pos_embeddings`` unchanged. + projection: Matrix A used to project Gaussian noise to the positional + embedding dimension. + noise_dim: Dimension of the sampled Gaussian noise before projection. + noise_std: Standard deviation of the Gaussian noise. If ``0.0``, + returns ``pos_embeddings`` unchanged. Returns: - Positional embeddings with optional gaussian noise added. + Positional embeddings with optional Gaussian noise added. """ - if not enabled or noise_std == 0.0: + if noise_std == 0.0: return pos_embeddings noise = torch.normal( diff --git a/tests/models/modules/test_ijepa_timm.py b/tests/models/modules/test_ijepa_timm.py index 6cc1b9ccf..a4cf8a550 100644 --- a/tests/models/modules/test_ijepa_timm.py +++ b/tests/models/modules/test_ijepa_timm.py @@ -1,9 +1,6 @@ -import unittest - import pytest import torch -from lightly.models import utils from lightly.utils import dependency if not dependency.timm_vit_available(): @@ -15,9 +12,8 @@ class TestIJEPAPredictorTIMM: - @pytest.mark.parametrize("use_stop", [True, False]) @pytest.mark.parametrize("noise_std", [0.0, 0.1]) - def test_init(self, use_stop: bool, noise_std: float) -> None: + def test_init(self, noise_std: float) -> None: IJEPAPredictorTIMM( num_patches=196, depth=2, @@ -28,14 +24,12 @@ def test_init(self, use_stop: bool, noise_std: float) -> None: mlp_ratio=4.0, proj_drop_rate=0.0, attn_drop_rate=0.0, - use_stop=use_stop, noise_std=noise_std, ) def _test_forward( self, device: torch.device, - use_stop: bool, noise_std: float, batch_size: int = 4, seed: int = 0, @@ -57,7 +51,6 @@ def _test_forward( mlp_ratio=4.0, proj_drop_rate=0.0, attn_drop_rate=0.0, - use_stop=use_stop, noise_std=noise_std, ).to(device) @@ -70,13 +63,11 @@ def _test_forward( assert list(predictions.shape) == [batch_size, num_patches, mlp_dim] assert torch.all(torch.isfinite(predictions)) - @pytest.mark.parametrize("use_stop", [True, False]) @pytest.mark.parametrize("noise_std", [0.0, 0.1]) - def test_forward(self, use_stop: bool, noise_std: float) -> None: - self._test_forward(torch.device("cpu"), use_stop, noise_std) + def test_forward(self, noise_std: float) -> None: + self._test_forward(torch.device("cpu"), noise_std) @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA not available.") - @pytest.mark.parametrize("use_stop", [True, False]) @pytest.mark.parametrize("noise_std", [0.0, 0.1]) - def test_forward_cuda(self, use_stop: bool, noise_std: float) -> None: - self._test_forward(torch.device("cuda"), use_stop, noise_std) + def test_forward_cuda(self, noise_std: float) -> None: + self._test_forward(torch.device("cuda"), noise_std) diff --git a/tests/utils/test_stochastic_positional_embedding.py b/tests/utils/test_stochastic_positional_embedding.py index 19bfe0958..ee9d686bd 100644 --- a/tests/utils/test_stochastic_positional_embedding.py +++ b/tests/utils/test_stochastic_positional_embedding.py @@ -8,16 +8,13 @@ def test_add_stochastic_positional_noise_disabled() -> None: pos_embeddings = torch.randn(2, 3, 4) out = utils.add_stochastic_positional_noise( - pos_embeddings=pos_embeddings, - projection=projection, - noise_dim=8, - enabled=False, + pos_embeddings=pos_embeddings, projection=projection, noise_dim=8, noise_std=0.0 ) assert torch.equal(out, pos_embeddings) -def test_add_stochastic_positional_noise_enabled_shape() -> None: +def test_add_stochastic_positional_noise_enabled() -> None: projection = torch.nn.Linear(8, 4) pos_embeddings = torch.randn(2, 3, 4) @@ -25,7 +22,8 @@ def test_add_stochastic_positional_noise_enabled_shape() -> None: pos_embeddings=pos_embeddings, projection=projection, noise_dim=8, - enabled=True, + noise_std=0.25, ) assert out.shape == pos_embeddings.shape + assert not torch.equal(out, pos_embeddings) From 14f29f85b4c84a01684558ef60e95ae31100ff96 Mon Sep 17 00:00:00 2001 From: kir_7 Date: Thu, 9 Jul 2026 23:50:49 +0530 Subject: [PATCH 4/7] Stop: fix typo --- lightly/models/modules/ijepa.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/lightly/models/modules/ijepa.py b/lightly/models/modules/ijepa.py index 2cd6f11b9..ce01caad5 100644 --- a/lightly/models/modules/ijepa.py +++ b/lightly/models/modules/ijepa.py @@ -184,7 +184,7 @@ class IJEPAEncoder(vision_transformer.Encoder): Encodes patch embeddings. Code inspired by [1]. - - [`0]: Joint-Embedding Predictive Architecture, 2023, https://arxiv.org/abs/2301.08243 + - [0]: Joint-Embedding Predictive Architecture, 2023, https://arxiv.org/abs/2301.08243 - [1]: https://github.com/facebookresearch/ijepa Attributes: From 345e9acbe35e0dd6c1df91e447ae888f03a280ba Mon Sep 17 00:00:00 2001 From: kir_7 Date: Mon, 13 Jul 2026 23:31:26 +0530 Subject: [PATCH 5/7] StoP: remove bias term in noise calculation --- lightly/models/modules/ijepa.py | 4 ++-- lightly/models/modules/ijepa_timm.py | 4 ++-- lightly/models/utils.py | 8 ++++---- tests/utils/test_stochastic_positional_embedding.py | 7 +++++-- 4 files changed, 13 insertions(+), 10 deletions(-) diff --git a/lightly/models/modules/ijepa.py b/lightly/models/modules/ijepa.py index ce01caad5..91a0110e7 100644 --- a/lightly/models/modules/ijepa.py +++ b/lightly/models/modules/ijepa.py @@ -153,10 +153,10 @@ def forward(self, x, masks_x, masks): pos_embs = utils.repeat_interleave_batch(pos_embs, B, repeat=len(masks_x)) # we add the stochastic positional embedding here: - # use self.predictor_embed as the projector + # use self.predictor_embed.weight as the projection matrix pos_embs = utils.add_stochastic_positional_noise( pos_embs, - self.predictor_embed, + self.predictor_embed.weight, noise_dim, noise_std=self.noise_std, ) diff --git a/lightly/models/modules/ijepa_timm.py b/lightly/models/modules/ijepa_timm.py index 9bd840fef..278a17d2a 100644 --- a/lightly/models/modules/ijepa_timm.py +++ b/lightly/models/modules/ijepa_timm.py @@ -142,10 +142,10 @@ def forward( pos_embs = utils.repeat_interleave_batch(pos_embs, B, repeat=len_masks_x) # we add the stochastic positional embedding here: - # use self.predictor_embed as the projector + # use self.predictor_embed.weight as the projection matrix pos_embs = utils.add_stochastic_positional_noise( pos_embs, - self.predictor_embed, + self.predictor_embed.weight, noise_dim, noise_std=self.noise_std, ) diff --git a/lightly/models/utils.py b/lightly/models/utils.py index b4877512d..aa9f9874d 100644 --- a/lightly/models/utils.py +++ b/lightly/models/utils.py @@ -1453,7 +1453,7 @@ def apply_masks(x: Tensor, masks: Tensor | list[Tensor]) -> Tensor: def add_stochastic_positional_noise( pos_embeddings: Tensor, - projection: Module, + projection_weight: Tensor, noise_dim: int, noise_std: float = 0.0, ) -> Tensor: @@ -1465,8 +1465,8 @@ def add_stochastic_positional_noise( Args: pos_embeddings: Positional embeddings of shape ``(batch_size, num_tokens, predictor_embed_dim)``. - projection: Matrix A used to project Gaussian noise to the positional - embedding dimension. + projection_weight: Matrix A used to project Gaussian noise to the positional + embedding dimension. Must have shape ``(predictor_embed_dim, noise_dim)``. noise_dim: Dimension of the sampled Gaussian noise before projection. noise_std: Standard deviation of the Gaussian noise. If ``0.0``, returns ``pos_embeddings`` unchanged. @@ -1485,4 +1485,4 @@ def add_stochastic_positional_noise( dtype=pos_embeddings.dtype, ) - return pos_embeddings + projection(noise) + return pos_embeddings + nn.functional.linear(noise, projection_weight, bias=None) diff --git a/tests/utils/test_stochastic_positional_embedding.py b/tests/utils/test_stochastic_positional_embedding.py index ee9d686bd..0fd8b4252 100644 --- a/tests/utils/test_stochastic_positional_embedding.py +++ b/tests/utils/test_stochastic_positional_embedding.py @@ -8,7 +8,10 @@ def test_add_stochastic_positional_noise_disabled() -> None: pos_embeddings = torch.randn(2, 3, 4) out = utils.add_stochastic_positional_noise( - pos_embeddings=pos_embeddings, projection=projection, noise_dim=8, noise_std=0.0 + pos_embeddings=pos_embeddings, + projection_weight=projection.weight, + noise_dim=8, + noise_std=0.0, ) assert torch.equal(out, pos_embeddings) @@ -20,7 +23,7 @@ def test_add_stochastic_positional_noise_enabled() -> None: out = utils.add_stochastic_positional_noise( pos_embeddings=pos_embeddings, - projection=projection, + projection_weight=projection.weight, noise_dim=8, noise_std=0.25, ) From ae95ecb0dfc5c151bfeade70f4bea2a8819ff062 Mon Sep 17 00:00:00 2001 From: kir_7 Date: Fri, 7 Aug 2026 19:04:50 +0530 Subject: [PATCH 6/7] StoP: disabling stop during eval --- lightly/models/modules/ijepa.py | 2 +- lightly/models/modules/ijepa_timm.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/lightly/models/modules/ijepa.py b/lightly/models/modules/ijepa.py index 91a0110e7..cac7b47e9 100644 --- a/lightly/models/modules/ijepa.py +++ b/lightly/models/modules/ijepa.py @@ -158,7 +158,7 @@ def forward(self, x, masks_x, masks): pos_embs, self.predictor_embed.weight, noise_dim, - noise_std=self.noise_std, + noise_std=self.noise_std if self.training else 0.0 ) pred_tokens = self.mask_token.repeat(pos_embs.size(0), pos_embs.size(1), 1) diff --git a/lightly/models/modules/ijepa_timm.py b/lightly/models/modules/ijepa_timm.py index 31e31911e..e449ca35e 100644 --- a/lightly/models/modules/ijepa_timm.py +++ b/lightly/models/modules/ijepa_timm.py @@ -153,7 +153,7 @@ def forward( pos_embs, self.predictor_embed.weight, noise_dim, - noise_std=self.noise_std, + noise_std=self.noise_std if self.training else 0.0 ) pred_tokens = self.decoder.mask_token.repeat( From 8e7655384e5eeff78551f9fdd702c0cca42c4d02 Mon Sep 17 00:00:00 2001 From: kir_7 Date: Sat, 8 Aug 2026 09:54:38 +0530 Subject: [PATCH 7/7] StoP: fixed formatting --- lightly/models/modules/ijepa.py | 2 +- lightly/models/modules/ijepa_timm.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/lightly/models/modules/ijepa.py b/lightly/models/modules/ijepa.py index cac7b47e9..6332df79a 100644 --- a/lightly/models/modules/ijepa.py +++ b/lightly/models/modules/ijepa.py @@ -158,7 +158,7 @@ def forward(self, x, masks_x, masks): pos_embs, self.predictor_embed.weight, noise_dim, - noise_std=self.noise_std if self.training else 0.0 + noise_std=self.noise_std if self.training else 0.0, ) pred_tokens = self.mask_token.repeat(pos_embs.size(0), pos_embs.size(1), 1) diff --git a/lightly/models/modules/ijepa_timm.py b/lightly/models/modules/ijepa_timm.py index e449ca35e..962de87df 100644 --- a/lightly/models/modules/ijepa_timm.py +++ b/lightly/models/modules/ijepa_timm.py @@ -153,7 +153,7 @@ def forward( pos_embs, self.predictor_embed.weight, noise_dim, - noise_std=self.noise_std if self.training else 0.0 + noise_std=self.noise_std if self.training else 0.0, ) pred_tokens = self.decoder.mask_token.repeat(