SIGN IN SIGN UP

[ENH] Add LeJEPA core building blocks (#1916)

* Add LeJEPALoss and invariance loss

Combine the existing SIGReg regularizer with a view-invariance term
to implement the full LeJEPA objective.

LeJEPALoss is a single-hyperparameter convex combination
(lambda * SIGReg + (1 - lambda) * invariance) with default
lambda_param=0.02 from the reference implementation.
SIGReg's constructor parameters are re-exposed for flexibility.
The invariance helper returns the mean squared deviation
of each view from the per-sample centroid across views.

Export LeJEPALoss from lightly.loss,
list both new objects in the loss rst page,
and add tests mirroring the SIGReg style.

* Add LeJEPAProjectionHead

Introduce the projection head used in LeJEPA as a 3-layer MLP
with BatchNorm and ReLU.

The defaults (input 2048, hidden 2048, output 256, 3 layers)
mirror the reference implementation's projector.
The constructor signature matches VICRegProjectionHead to stay
consistent with the library's existing head conventions.

Register LeJEPAProjectionHead in the projection-head test suite
so the existing test_single_projection_head loop covers it
across the standard set of dimension combinations.

* Add LeJEPAEncoder module

Introduce LeJEPAEncoder as a convenience module that wraps
a backbone and a LeJEPAProjectionHead so a single forward call
maps input images to projected embeddings.

The encoder flattens backbone outputs from dim=1 onward
before passing them to the projection head, matching the
pattern used by the LeJEPA examples.

Export LeJEPAProjectionHead and LeJEPAEncoder from the
lightly.models.modules package so users can import them directly,
list the new module in the rst docs,
and add tests covering forward shape and gradient flow.

* Relocate LeJEPALoss to lightly.models.modules.lejepa

Move LeJEPALoss and lejepa_invariance_loss out of
lightly/loss/lejepa_loss.py and into
lightly/models/modules/lejepa.py to match the structure
specified in issue comment at https://github.com/lightly-ai/lightly/issues/1888#issuecomment-3996087102

The new location groups all method-specific LeJEPA components
(the encoder wrapper, the invariance formula, and the combined
objective) in one file, while SIGReg stays in lightly.loss
as a reusable isotropic-Gaussian regularizer usable on its own.

Update the loss __init__ to drop the LeJEPALoss re-export,
expose LeJEPALoss via lightly.models.modules,
move the corresponding test classes to tests/models/modules/test_lejepa.py,
and update lightly.loss.rst (LeJEPALoss is now documented
through the lightly.models.rst .lejepa automodule section).

* Add LeJEPA wiring and projector backward tests

The wiring tests verify that LeJEPALoss at lambda=0 collapses
to the standalone invariance loss, and at lambda=1 collapses
to a standalone SIGReg call under the same random seed.
Without these checks, a bug that dropped either term would still
pass the earlier forward/backward/distributed tests.

The projector backward test exercises LeJEPAProjectionHead across
the same dimension combinations as the forward sweep,
asserts the expected output shape,
and confirms that gradients exist with the correct shape on the
input and on every named parameter after a backward pass.

* Adjust LeJEPALoss docstring bullet indent for readability

* fix: address lejepa review comments

* assert grads are not none

---------

Co-authored-by: gabrielfruet <gabrielfruet538@gmail.com>
N
Nirbhai committed
5aa9d51daeedf44a6974d83c4d5c14c91d803ede
Parent: 753643f
Committed by GitHub <noreply@github.com> on 5/18/2026, 7:06:47 AM