From 394b3cd153f31c180a70694ed41acb1d7a0ba41e Mon Sep 17 00:00:00 2001 From: zephyrus Date: Sun, 2 Aug 2026 02:25:50 +0300 Subject: [PATCH] test(vision-encoder): add coverage for optimizer weight decay and filtering --- .../tests/test_vision_encoder.py | 16 +++++++++++++++- 1 file changed, 15 insertions(+), 1 deletion(-) diff --git a/contrastive-pretraining/tests/test_vision_encoder.py b/contrastive-pretraining/tests/test_vision_encoder.py index 316f7b9..16f866b 100644 --- a/contrastive-pretraining/tests/test_vision_encoder.py +++ b/contrastive-pretraining/tests/test_vision_encoder.py @@ -3,6 +3,8 @@ import torch import torch.nn as nn from vision_encoder.vjepa_encoder import ResidualTemporalDownsample +from vision_encoder.optimizer import get_optimizer + class TestResidualTemporalDownsample: @@ -143,8 +145,20 @@ class TestOptimizerImport: """Test that the optimizer utility from vision_encoder works.""" def test_get_optimizer_creates_adam(self): - from vision_encoder.optimizer import get_optimizer params = nn.Linear(10, 10).parameters() opt = get_optimizer(set(params), lr=1e-3, wd=0.01) assert opt is not None assert len(opt.param_groups) > 0 + + def test_get_optimizer_zero_weight_decay(self): + from torch.optim import Adam + params = nn.Linear(10, 10).parameters() + opt = get_optimizer(set(params), lr=1e-3, wd=0) + assert isinstance(opt, Adam) + + def test_get_optimizer_filter_by_requires_grad(self): + p1 = nn.Parameter(torch.randn(10, 10), requires_grad=True) + p2 = nn.Parameter(torch.randn(10, 10), requires_grad=False) + opt = get_optimizer(params=[p1, p2], lr=1e-3, filter_by_requires_grad=True) + all_params = [p for group in opt.param_groups for p in group['params']] + assert all_params == [p1]