From 537c0dc98dc470f73bef13bc8babb5edd6f9a11a Mon Sep 17 00:00:00 2001 From: lezcano Date: Sun, 7 Jun 2026 11:16:43 +0100 Subject: [PATCH] Reduce manifold sampling temporaries --- geotorch/fixedrank.py | 4 ++-- geotorch/lowrank.py | 2 +- geotorch/pssdfixedrank.py | 2 +- geotorch/sphere.py | 5 ++--- geotorch/symmetric.py | 2 +- test/test_lowrank.py | 9 +++++++++ test/test_positive_semidefinite.py | 9 +++++++++ test/test_sphere.py | 8 ++++++++ 8 files changed, 33 insertions(+), 8 deletions(-) diff --git a/geotorch/fixedrank.py b/geotorch/fixedrank.py index 70d3e511..400e030f 100644 --- a/geotorch/fixedrank.py +++ b/geotorch/fixedrank.py @@ -113,8 +113,8 @@ def sample(self, init_=torch.nn.init.xavier_normal_, eps=5e-6, factorized=False) """ U, S, V = super().sample(factorized=True, init_=init_) with torch.no_grad(): - # S >= 0, as given by torch.linalg.eigvalsh() - S[S < eps] = eps + # S >= 0, as given by torch.linalg.svd() + S.clamp_min_(eps) if factorized: return U, S, V else: diff --git a/geotorch/lowrank.py b/geotorch/lowrank.py index 36b7f703..5c5d4320 100644 --- a/geotorch/lowrank.py +++ b/geotorch/lowrank.py @@ -107,7 +107,7 @@ def in_manifold_singular_values(self, S, eps=1e-5): return True # We compute the \infty-norm of the remaining dimension D = S[..., self.rank :] - infty_norm_err = D.abs().max(dim=-1).values + infty_norm_err = D.abs().amax(dim=-1) return (infty_norm_err < eps).all() def in_manifold(self, X, eps=1e-5): diff --git a/geotorch/pssdfixedrank.py b/geotorch/pssdfixedrank.py index 62211b13..3af1997a 100644 --- a/geotorch/pssdfixedrank.py +++ b/geotorch/pssdfixedrank.py @@ -95,5 +95,5 @@ def sample(self, init_=torch.nn.init.xavier_normal_, eps=5e-6): L, Q = super().sample(factorized=True, init_=init_) with torch.no_grad(): # L >= 0, as given by torch.linalg.eigvalsh() - L[L < eps] = eps + L.clamp_min_(eps) return (Q * L.unsqueeze(-2)) @ Q.transpose(-2, -1) diff --git a/geotorch/sphere.py b/geotorch/sphere.py index 71de4cb3..ad1ed534 100644 --- a/geotorch/sphere.py +++ b/geotorch/sphere.py @@ -16,14 +16,13 @@ def uniform_init_sphere_(x, r=1.0): """ with torch.no_grad(): x.normal_() - x.copy_(r * project(x)) + x.copy_(project(x)).mul_(r) return x def _in_sphere(x, r, eps): norm = torch.linalg.vector_norm(x, dim=-1) - rs = torch.full_like(norm, r) - return (torch.linalg.vector_norm(norm - rs, ord=float("inf")) < eps).all() + return (norm - r).abs().amax() < eps class SphereEmbedded(nn.Module): diff --git a/geotorch/symmetric.py b/geotorch/symmetric.py index 74fa34fc..4b2aba48 100644 --- a/geotorch/symmetric.py +++ b/geotorch/symmetric.py @@ -168,7 +168,7 @@ def in_manifold_eigen(self, L, eps=1e-6): if L.size(-1) > self.rank: # We compute the \infty-norm of the remaining dimension D = L[..., : -self.rank] - infty_norm_err = D.abs().max(dim=-1).values + infty_norm_err = D.abs().amax(dim=-1) if (infty_norm_err > 5.0 * eps).any(): return False return (L[..., -self.rank :] >= -eps).all().item() diff --git a/test/test_lowrank.py b/test/test_lowrank.py index 5ba890ab..8e193126 100644 --- a/test/test_lowrank.py +++ b/test/test_lowrank.py @@ -56,3 +56,12 @@ def init_(X): with mock.patch("torch.linalg.svd", side_effect=AssertionError): sample = manifold.sample(init_) self.assertTrue(torch.equal(sample, expected)) + + def test_fixed_rank_sample_clamps_singular_values(self): + manifold = FixedRank(size=(3, 3), rank=2).double() + _, singular_values, _ = manifold.sample( + init_=lambda X: X.zero_(), eps=0.25, factorized=True + ) + self.assertTrue( + torch.equal(singular_values, torch.full_like(singular_values, 0.25)) + ) diff --git a/test/test_positive_semidefinite.py b/test/test_positive_semidefinite.py index e56cc28b..5f58d17a 100644 --- a/test/test_positive_semidefinite.py +++ b/test/test_positive_semidefinite.py @@ -1,5 +1,7 @@ from unittest import TestCase +import torch + from geotorch.pssdlowrank import PSSDLowRank from geotorch.pssdfixedrank import PSSDFixedRank from geotorch.pssd import PSSD @@ -40,3 +42,10 @@ def test_positive_semidefinite_errors(self): PSD(size=(5, 2), f=3) with self.assertRaises(ValueError): PSD(size=(5, 3), f="fail") + + def test_fixed_rank_sample_clamps_eigenvalues(self): + manifold = PSSDFixedRank(size=(3, 3), rank=2).double() + sample = manifold.sample(init_=lambda X: X.zero_(), eps=0.25) + eigenvalues = torch.linalg.eigvalsh(sample) + expected = torch.tensor([0.0, 0.25, 0.25], dtype=torch.float64) + torch.testing.assert_close(eigenvalues, expected) diff --git a/test/test_sphere.py b/test/test_sphere.py index 5201fb7d..803a95ab 100644 --- a/test/test_sphere.py +++ b/test/test_sphere.py @@ -51,3 +51,11 @@ def test_embedded_sample_uses_module_dtype(self): sample = sphere.sample() self.assertEqual(sample.dtype, torch.float64) self.assertEqual(sample.shape, (2, 3)) + + def test_membership_preserves_strict_epsilon_boundary(self): + sphere = SphereEmbedded(size=(1,), radius=2.0) + x = torch.tensor([2.0001], dtype=torch.float64) + eps = (torch.linalg.vector_norm(x) - sphere.radius).item() + + self.assertFalse(sphere.in_manifold(x, eps=eps)) + self.assertTrue(sphere.in_manifold(x, eps=2.0 * eps))