Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 11 additions & 4 deletions src/yetanotherspdnet/nn/batchnorm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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}")
Comment on lines +359 to +360
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:
Expand Down Expand Up @@ -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

Expand Down
32 changes: 32 additions & 0 deletions tests/nn/test_batchnorm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Loading