From 0332435f772962f07a019b2080ac32f0938b75dc Mon Sep 17 00:00:00 2001 From: Matthieu Gallet Date: Tue, 29 Sep 2026 11:40:44 +0200 Subject: [PATCH] feat: configurable initial t for the adaptive GAH batch norm The learnable t of the adaptive GAH mean was hard-coded to 0.5, so the t_init options of downstream configs (spdnet-training batchnorm_t_gah_init) had no effect. mean_options={"t_init": t} now sets it (validated in [0, 1]; default 0.5 unchanged). Docstrings list the adaptive GAH mean type. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_01PQdVCDbXCd8gvf1Y4TufJR --- src/yetanotherspdnet/nn/batchnorm.py | 15 +++++++++---- tests/nn/test_batchnorm.py | 32 ++++++++++++++++++++++++++++ 2 files changed, 43 insertions(+), 4 deletions(-) diff --git a/src/yetanotherspdnet/nn/batchnorm.py b/src/yetanotherspdnet/nn/batchnorm.py index 1fb1d29..b8a59cb 100644 --- a/src/yetanotherspdnet/nn/batchnorm.py +++ b/src/yetanotherspdnet/nn/batchnorm.py @@ -159,12 +159,14 @@ def __init__( Choice of SPD mean. Default is "affine_invariant". Choices are: "affine_invariant", "log_euclidean", "arithmetic", "harmonic", "geometric_arithmetic_harmonic", - "bures_wasserstein" + "adaptive_geometric_arithmetic_harmonic", "bures_wasserstein" mean_options : dict | None, optional Options for the SPD mean computation. For affine-invariant mean, one can typically set {'n_iterations': 5}. - Currently, for others, no options available. + For adaptive_geometric_arithmetic_harmonic, {'t_init': t} sets the + initial value of the learnable t in [0, 1] (0: harmonic, + 1: arithmetic). Default t_init is 0.5 (the GAH mean). Default is None momentum : float, optional @@ -353,8 +355,11 @@ def _init_mean(self) -> None: elif self.mean_type == "adaptive_geometric_arithmetic_harmonic": # Register learnable parameter t for adaptive GAH mean # Use sigmoid parametrization to constrain t in [0, 1] + t_init = (self.mean_options or {}).get("t_init", 0.5) + if not 0.0 <= t_init <= 1.0: + raise ValueError(f"t_init must lie in [0, 1], got {t_init}") self.t_gah = torch.nn.Parameter( - torch.tensor(0.5, dtype=self.dtype, device=self.device) + torch.tensor(t_init, dtype=self.dtype, device=self.device) ) register_parametrization(self, "t_gah", ScalarSigmoidParametrization()) if self.use_autograd: @@ -679,11 +684,13 @@ def __init__( Choice of SPD mean. Default is "affine_invariant". Choices are: "affine_invariant", "log_euclidean", "arithmetic", "harmonic", "geometric_arithmetic_harmonic", - "bures_wasserstein" + "adaptive_geometric_arithmetic_harmonic", "bures_wasserstein" mean_options : dict | None, optional Options for the SPD mean computation. For affine-invariant mean, one can typically set {'n_iterations': 5}. + For adaptive_geometric_arithmetic_harmonic, one can set + {'t_init': 0.5} (initial learnable t, 0: harmonic, 1: arithmetic). For bures_wasserstein, one can set {'n_iterations': 1}. Default is None diff --git a/tests/nn/test_batchnorm.py b/tests/nn/test_batchnorm.py index 3e8b80e..4859385 100644 --- a/tests/nn/test_batchnorm.py +++ b/tests/nn/test_batchnorm.py @@ -2229,3 +2229,35 @@ def test_models_forward_bw_options(self, bw_theta): m.bw_theta == bw_theta and m.bw_batch_stats_grad is False for m in layers ) + + +class TestAdaptiveGAHInit: + """``mean_options={'t_init': t}`` sets the initial learnable t of the + adaptive GAH mean (sigmoid-parametrized); the default stays 0.5.""" + + @pytest.mark.parametrize("t_init", [None, 0.2, 0.9]) + def test_t_init(self, t_init, device, dtype, generator): + options = None if t_init is None else {"t_init": t_init} + layer = batchnorm.BatchNormSPDMean( + 5, + mean_type="adaptive_geometric_arithmetic_harmonic", + mean_options=options, + device=device, + dtype=dtype, + ) + expected = 0.5 if t_init is None else t_init + assert_close(layer.t_gah, torch.tensor(expected, device=device, dtype=dtype)) + data = random_SPD(5, 8, device=device, dtype=dtype, generator=generator) + layer(data).sum().backward() + original = layer.parametrizations.t_gah.original + assert original.grad is not None and torch.isfinite(original.grad) + + def test_t_init_out_of_range(self, device, dtype): + with pytest.raises(ValueError, match="t_init"): + batchnorm.BatchNormSPDMean( + 5, + mean_type="adaptive_geometric_arithmetic_harmonic", + mean_options={"t_init": 1.5}, + device=device, + dtype=dtype, + )