[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