From 4dcdc6db6b884021d8da9ee7bf0d070ca59d4ec5 Mon Sep 17 00:00:00 2001 From: dara <59355936+DaraVaram@users.noreply.github.com> Date: Tue, 29 Sep 2026 02:17:00 +0400 Subject: [PATCH 1/5] feat(aggregation): Add PCD * Add PCD and PCDWeighting (Priority-Constrained Descent, TMLR 2026) * Solve the QP with a Gramian-only Goldfarb-Idnani dual active-set method * Add tests, docs page, README entry and changelog entry Co-Authored-By: Claude Opus 5.5 --- CHANGELOG.md | 9 + README.md | 1 + docs/source/docs/aggregation/index.rst | 1 + docs/source/docs/aggregation/pcd.rst | 10 + src/torchjd/aggregation/__init__.py | 3 + src/torchjd/aggregation/_pcd.py | 321 +++++++++++++++++++++++++ tests/unit/aggregation/test_pcd.py | 242 +++++++++++++++++++ tests/unit/aggregation/test_values.py | 4 + 8 files changed, 591 insertions(+) create mode 100644 docs/source/docs/aggregation/pcd.rst create mode 100644 src/torchjd/aggregation/_pcd.py create mode 100644 tests/unit/aggregation/test_pcd.py diff --git a/CHANGELOG.md b/CHANGELOG.md index 69116c24d..40e0d416c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,6 +8,15 @@ changelog does not include internal changes that do not affect the user. ## [Unreleased] +### Added + +- Added `PCD` and `PCDWeighting` from [Not All Objectives Are Born Equal: Priority-Constrained + Descent for Hierarchical Multi-Objective Optimization](https://openreview.net/forum?id=HT01yGHLEt) + (TMLR 2026). The first row of the Jacobian is treated as the primary objective: `PCD` follows its + gradient as closely as possible, subject to each other objective receiving at least a fraction + `tau` of normalized first-order progress. `PCD` and `PCDWeighting` are stateful: they normalize + the gradients by a bias-corrected moving average of their squared norms. + ## [0.17.1] - 2026-09-23 ### Fixed diff --git a/README.md b/README.md index ff862bc4c..9d53149fd 100644 --- a/README.md +++ b/README.md @@ -192,6 +192,7 @@ TorchJD provides many existing aggregators from the literature, listed in the fo | [MGDA](https://torchjd.org/stable/docs/aggregation/mgda#torchjd.aggregation.MGDA) | [MGDAWeighting](https://torchjd.org/stable/docs/aggregation/mgda#torchjd.aggregation.MGDAWeighting) | [Multiple-gradient descent algorithm (MGDA) for multiobjective optimization](https://comptes-rendus.academie-sciences.fr/mathematique/articles/10.1016/j.crma.2012.03.014/) | | - | [MoDoWeighting](https://torchjd.org/stable/docs/aggregation/modo/#torchjd.aggregation.MoDoWeighting) | [Three-Way Trade-Off in Multi-Objective Learning: Optimization, Generalization and Conflict-Avoidance](https://www.jmlr.org/papers/volume25/23-1287/23-1287.pdf) | | [NashMTL](https://torchjd.org/stable/docs/aggregation/nash_mtl#torchjd.aggregation.NashMTL) | - | [Multi-Task Learning as a Bargaining Game](https://arxiv.org/pdf/2202.01017) | +| [PCD](https://torchjd.org/stable/docs/aggregation/pcd#torchjd.aggregation.PCD) | [PCDWeighting](https://torchjd.org/stable/docs/aggregation/pcd#torchjd.aggregation.PCDWeighting) | [Not All Objectives Are Born Equal: Priority-Constrained Descent for Hierarchical Multi-Objective Optimization](https://openreview.net/forum?id=HT01yGHLEt) | | [PCGrad](https://torchjd.org/stable/docs/aggregation/pcgrad#torchjd.aggregation.PCGrad) | [PCGradWeighting](https://torchjd.org/stable/docs/aggregation/pcgrad#torchjd.aggregation.PCGradWeighting) | [Gradient Surgery for Multi-Task Learning](https://arxiv.org/pdf/2001.06782) | | [Random](https://torchjd.org/stable/docs/aggregation/random#torchjd.aggregation.Random) | [RandomWeighting](https://torchjd.org/stable/docs/aggregation/random#torchjd.aggregation.RandomWeighting) | [Reasonable Effectiveness of Random Weighting: A Litmus Test for Multi-Task Learning](https://arxiv.org/pdf/2111.10603) | | - | [SDMGradWeighting](https://torchjd.org/stable/docs/aggregation/sdmgrad#torchjd.aggregation.SDMGradWeighting) | [Direction-oriented Multi-objective Learning: Simple and Provable Stochastic Algorithms](https://arxiv.org/pdf/2305.18409) | diff --git a/docs/source/docs/aggregation/index.rst b/docs/source/docs/aggregation/index.rst index b77332dbf..cb4cc77a7 100644 --- a/docs/source/docs/aggregation/index.rst +++ b/docs/source/docs/aggregation/index.rst @@ -40,6 +40,7 @@ Abstract base classes mgda.rst modo.rst nash_mtl.rst + pcd.rst pcgrad.rst random.rst sdmgrad.rst diff --git a/docs/source/docs/aggregation/pcd.rst b/docs/source/docs/aggregation/pcd.rst new file mode 100644 index 000000000..cbbb8269b --- /dev/null +++ b/docs/source/docs/aggregation/pcd.rst @@ -0,0 +1,10 @@ +:hide-toc: + +PCD +=== + +.. autoclass:: torchjd.aggregation.PCD + :members: __call__, reset + +.. autoclass:: torchjd.aggregation.PCDWeighting + :members: __call__, reset diff --git a/src/torchjd/aggregation/__init__.py b/src/torchjd/aggregation/__init__.py index 55e855eed..826cbc04f 100644 --- a/src/torchjd/aggregation/__init__.py +++ b/src/torchjd/aggregation/__init__.py @@ -55,6 +55,7 @@ from ._mgda import MGDA, MGDAWeighting from ._modo import MoDoWeighting from ._nash_mtl import NashMTL +from ._pcd import PCD, PCDWeighting from ._pcgrad import PCGrad, PCGradWeighting from ._random import Random, RandomWeighting from ._sdmgrad import SDMGradWeighting @@ -93,6 +94,8 @@ "MGDAWeighting", "MoDoWeighting", "NashMTL", + "PCD", + "PCDWeighting", "PCGrad", "PCGradWeighting", "Random", diff --git a/src/torchjd/aggregation/_pcd.py b/src/torchjd/aggregation/_pcd.py new file mode 100644 index 000000000..5e11893d9 --- /dev/null +++ b/src/torchjd/aggregation/_pcd.py @@ -0,0 +1,321 @@ +from __future__ import annotations + +from typing import cast + +import torch +from torch import Tensor + +from torchjd._mixins import Stateful +from torchjd.linalg import PSDMatrix + +from ._aggregator_bases import GramianWeightedAggregator +from ._mixins import _NonDifferentiable +from ._weighting_bases import _GramianWeighting + + +# Non-differentiable: the weights are obtained from an active-set solver, which branches on the +# values of the Gramian. +class PCDWeighting(_GramianWeighting, Stateful, _NonDifferentiable): + r""" + :class:`~torchjd.Stateful` + :class:`~torchjd.aggregation.Weighting` [:class:`~torchjd.linalg.PSDMatrix`] + giving the weights of :class:`~torchjd.aggregation.PCD`. + + The first row and column of the Gramian correspond to the primary objective. + + :param tau: The fraction :math:`\tau \in [0, 1]` of normalized first-order progress guaranteed + to each secondary objective. Either a float, shared by all secondary objectives, or a vector + with one value per secondary objective (i.e. of length :math:`m - 1`). + :param beta: The decay :math:`\beta \in [0, 1)` of the exponential moving average of the + squared gradient norms. + :param eps: The non-negative constant :math:`\epsilon` added to the moving average before + taking its inverse square root. + + .. note:: + The quadratic program is solved exactly with the dual active-set method of `Goldfarb and + Idnani (1983) `_, expressed in terms of the Gramian + only. The reference implementation instead enumerates the working sets of constraints + (Appendix B.2 of the paper). Both methods find the same minimizer, but the enumeration is + exponential in :math:`m`. + """ + + def __init__(self, tau: float | Tensor = 0.02, beta: float = 0.999, eps: float = 1e-8) -> None: + super().__init__() + self.tau = tau + self.beta = beta + self.eps = eps + self.register_buffer("_sq_norm_ema", None) + self.register_buffer("_n_steps", None) + + @property + def tau(self) -> float | Tensor: + return self._tau + + @tau.setter + def tau(self, value: float | Tensor) -> None: + if isinstance(value, Tensor): + if value.ndim != 1: + raise ValueError( + f"Attribute `tau` must be a float or a vector (1D Tensor). Found `tau.ndim = " + f"{value.ndim}`.", + ) + is_valid = bool(((value >= 0.0) & (value <= 1.0)).all()) + else: + is_valid = 0.0 <= value <= 1.0 + if not is_valid: + raise ValueError(f"Attribute `tau` must be in [0, 1]. Found tau={value!r}.") + self._tau = value + + @property + def beta(self) -> float: + return self._beta + + @beta.setter + def beta(self, value: float) -> None: + if not (0.0 <= value < 1.0): + raise ValueError(f"Attribute `beta` must be in [0, 1). Found beta={value!r}.") + self._beta = value + + @property + def eps(self) -> float: + return self._eps + + @eps.setter + def eps(self, value: float) -> None: + if not (value >= 0.0): + raise ValueError(f"Attribute `eps` must be non-negative. Found eps={value!r}.") + self._eps = value + + def reset(self) -> None: + """Clears the moving average of the squared gradient norms.""" + + self._sq_norm_ema = None + self._n_steps = None + + def forward(self, gramian: PSDMatrix, /) -> Tensor: + # The problem only has size m x m, so we solve it on cpu and in float64. + G = gramian.to(device="cpu", dtype=torch.float64) + m = G.shape[0] + weights = torch.zeros(m, dtype=torch.float64) + if m == 0: + return weights.to(device=gramian.device, dtype=gramian.dtype) + + taus = self._get_taus(m) + sq_norms = G.diagonal() + scales = self._update_scales(sq_norms) + + # If the gradient of the primary objective is zero, the update is zero. + if sq_norms[0] > 0.0: + normalized_G = G * torch.outer(scales, scales) + coefficients = _solve_qp(normalized_G, taus) * scales + direction_sq_norm = coefficients @ G @ coefficients + # Rescale the direction to the norm of the primary gradient. + if direction_sq_norm > 0.0: + weights = coefficients * (sq_norms[0] / direction_sq_norm).sqrt() + + return weights.to(device=gramian.device, dtype=gramian.dtype) + + def _get_taus(self, m: int) -> Tensor: + if isinstance(self.tau, Tensor): + if self.tau.shape[0] != m - 1: + raise ValueError( + "When `tau` is a vector, it must have one value per secondary objective. Found " + f"`tau.shape[0] = {self.tau.shape[0]}` for {m - 1} secondary objectives.", + ) + return self.tau.to(device="cpu", dtype=torch.float64) + return torch.full([m - 1], self.tau, dtype=torch.float64) + + def _update_scales(self, sq_norms: Tensor) -> Tensor: + """ + Updates the moving average of the squared gradient norms with the current ones, and returns + the scale by which each gradient should be multiplied to be normalized. + """ + + m = sq_norms.shape[0] + if self._sq_norm_ema is None or self._sq_norm_ema.shape[0] != m: + self._sq_norm_ema = torch.zeros(m, dtype=torch.float64) + self._n_steps = torch.tensor(0) + + sq_norm_ema = self._sq_norm_ema.to(sq_norms) + n_steps = int(cast(Tensor, self._n_steps)) + 1 + sq_norm_ema = self.beta * sq_norm_ema + (1.0 - self.beta) * sq_norms + self._sq_norm_ema = sq_norm_ema + self._n_steps = torch.tensor(n_steps) + + debiased_sq_norm_ema = sq_norm_ema / (1.0 - self.beta**n_steps) + return 1.0 / (debiased_sq_norm_ema + self.eps).sqrt() + + def __repr__(self) -> str: + return f"{self.__class__.__name__}(tau={self.tau!r}, beta={self.beta!r}, eps={self.eps!r})" + + +class PCD(GramianWeightedAggregator, Stateful, _NonDifferentiable): + r""" + :class:`~torchjd.Stateful` + :class:`~torchjd.aggregation.GramianWeightedAggregator` implementing Priority-Constrained + Descent (PCD), from `Not All Objectives Are Born Equal: Priority-Constrained Descent for + Hierarchical Multi-Objective Optimization `_ (TMLR + 2026, `arXiv:2606.29521 `_). + + The first row :math:`g_1` of the input matrix is the gradient of the primary objective, and the + other rows :math:`g_2, \dots, g_m` are the gradients of the secondary objectives. PCD follows the + primary gradient as closely as possible, subject to each secondary objective receiving at least + a fraction :math:`\tau` of the first-order progress that a step along its own normalized + gradient would give: + + .. math:: + + \tilde d = \mathop{\mathrm{arg\,min}}_{d} \ \frac{1}{2} \left\| d - \tilde g_1 \right\|^2 + \quad \text{subject to} \quad \tilde g_j^\top d \geq \tau \left\| \tilde g_j \right\|^2 + \quad \text{for } j = 2, \dots, m, + + where :math:`\tilde g_i = s_i g_i` is the normalized gradient of objective :math:`i`, with + :math:`s_i = 1 / \sqrt{\hat v_i + \epsilon}`, and :math:`\hat v_i` is the bias-corrected + exponential moving average, over the successive calls, of :math:`\|g_i\|^2`. The output is + :math:`\tilde d` rescaled to the norm of :math:`g_1`. Because of the normalization, :math:`\tau` + is a scale-free fraction. + + Special cases: + + - If :math:`g_1 = 0`, the output is zero. + - If the constraints cannot be satisfied simultaneously (which requires :math:`m \geq 3`, e.g. + two anti-parallel secondary gradients), they are dropped and the output is :math:`g_1`. + - If :math:`\tilde d = 0`, the output is zero. + + :param tau: The fraction :math:`\tau \in [0, 1]` of normalized first-order progress guaranteed + to each secondary objective. Either a float, shared by all secondary objectives, or a vector + with one value per secondary objective (i.e. of length :math:`m - 1`). + :param beta: The decay :math:`\beta \in [0, 1)` of the exponential moving average of the + squared gradient norms. + :param eps: The non-negative constant :math:`\epsilon` added to the moving average before + taking its inverse square root. + + .. note:: + PCD is not symmetric in the objectives: the first row of the input matrix (e.g. the + first loss given to :func:`~torchjd.autojac.backward`) is the primary objective. + + .. note:: + This aggregator is stateful: it keeps the moving average of the squared gradient norms + across calls. Use :meth:`reset` to clear it. It is also cleared automatically when the + number of rows changes. + + .. note:: + The reference implementation is available at + `github.com/DaraVaram/priority-constrained-descent + `_. + """ + + gramian_weighting: PCDWeighting + + def __init__(self, tau: float | Tensor = 0.02, beta: float = 0.999, eps: float = 1e-8) -> None: + super().__init__(PCDWeighting(tau=tau, beta=beta, eps=eps)) + + @property + def tau(self) -> float | Tensor: + return self.gramian_weighting.tau + + @tau.setter + def tau(self, value: float | Tensor) -> None: + self.gramian_weighting.tau = value + + @property + def beta(self) -> float: + return self.gramian_weighting.beta + + @beta.setter + def beta(self, value: float) -> None: + self.gramian_weighting.beta = value + + @property + def eps(self) -> float: + return self.gramian_weighting.eps + + @eps.setter + def eps(self, value: float) -> None: + self.gramian_weighting.eps = value + + def reset(self) -> None: + """Clears the moving average of the squared gradient norms.""" + + self.gramian_weighting.reset() + + def __repr__(self) -> str: + return f"{self.__class__.__name__}(tau={self.tau!r}, beta={self.beta!r}, eps={self.eps!r})" + + +def _solve_qp(G: Tensor, taus: Tensor, tol: float = 1e-9) -> Tensor: + r""" + Solves the quadratic program of PCD, given the Gramian :math:`G` of the (normalized) gradients + :math:`g_1, \dots, g_m`, whose first row and column correspond to the primary objective: + + .. math:: + + \min_d \frac{1}{2} \|d - g_1\|^2 \quad \text{s.t.} \quad + g_j^\top d \geq \tau_j \|g_j\|^2 \quad \text{for } j = 2, \dots, m + + Returns the weights :math:`w = [1, \mu_2, \dots, \mu_m]` such that the solution is + :math:`\sum_i w_i g_i`, where :math:`\mu_j \geq 0` are the KKT multipliers of the constraints. + If the constraints are infeasible, returns :math:`[1, 0, \dots, 0]`. + + This is the dual active-set method of Goldfarb and Idnani (1983), for an identity Hessian. The + stationarity condition :math:`d = g_1 + \sum_j \mu_j g_j` holds at every iteration, so that all + the required inner products can be read from :math:`G`. + """ + + m = G.shape[0] + weights = torch.zeros(m, dtype=G.dtype) + weights[0] = 1.0 + if m == 1: + return weights + + Gs = G[1:, 1:] + b = taus * Gs.diagonal() - G[1:, 0] # b_j > 0 iff g_1 violates the constraint of objective j + mu = torch.zeros(m - 1, dtype=G.dtype) + atol = tol * max(G.diagonal().max().item(), torch.finfo(G.dtype).tiny) + active: list[int] = [] + + while True: + slacks = Gs @ mu - b + slacks[active] = 0.0 # The active constraints hold with equality, up to rounding errors. + p = int(slacks.argmin()) + if slacks[p] >= -atol: + weights[1:] = mu.clamp(min=0.0) + return weights + + # Add the violated constraint p to the active set, possibly dropping some others first. + slack_p = slacks[p] + while True: + # In the dual space, the step direction is +1 for mu_p and -r for the active mu_j. + # In the primal space, it is z = g_p - sum_j r_j g_j, the component of g_p orthogonal to + # the gradients of the active constraints. + if len(active) > 0: + r = torch.linalg.solve(Gs[active][:, active], Gs[active, p]) + z_sq_norm = Gs[p, p] - Gs[p, active] @ r + ratios = torch.where(r > 0.0, mu[active] / r, torch.inf) + k = int(ratios.argmin()) + partial_step = ratios[k].item() + else: + r = torch.zeros(0, dtype=G.dtype) + z_sq_norm = Gs[p, p] + k = -1 + partial_step = torch.inf + + # If z = 0, g_p is a linear combination of the active gradients and no primal step is + # possible. + is_z_non_zero = z_sq_norm > tol * Gs[p, p] + full_step = (-slack_p / z_sq_norm).item() if is_z_non_zero else torch.inf + + step = min(partial_step, full_step) + if step == torch.inf: + return weights # The constraints are infeasible. + + mu[active] -= step * r + mu[p] += step + if full_step <= partial_step: + active.append(p) + break + + # Drop the constraint whose multiplier reached zero, and try adding p again. + mu[active[k]] = 0.0 + del active[k] + slack_p = Gs[p] @ mu - b[p] diff --git a/tests/unit/aggregation/test_pcd.py b/tests/unit/aggregation/test_pcd.py new file mode 100644 index 000000000..40929425d --- /dev/null +++ b/tests/unit/aggregation/test_pcd.py @@ -0,0 +1,242 @@ +import torch +from pytest import mark, raises +from torch import Tensor +from torch.testing import assert_close +from utils.tensors import ones_, randn_, tensor_, zeros_ + +from torchjd.aggregation import PCD, PCDWeighting +from torchjd.aggregation._pcd import _solve_qp + +from ._asserts import assert_expected_structure, assert_non_differentiable +from ._inputs import scaled_matrices, typical_matrices, typical_matrices_2_plus_rows + +scaled_pairs = [(PCD(), m) for m in scaled_matrices] +typical_pairs = [(PCD(), m) for m in typical_matrices] +requires_grad_pairs = [(PCD(), ones_(3, 5, requires_grad=True))] + + +def _two_objectives_pcd(matrix: Tensor, tau: float, scales: Tensor) -> Tensor: + """ + Closed-form solution of PCD with one secondary objective (Corollary 4.6 of the paper), used to + derive expected values independently of the implementation. + """ + + g1, g2 = matrix[0] * scales[0], matrix[1] * scales[1] + mu = ((tau * (g2 @ g2) - g1 @ g2) / (g2 @ g2)).clamp(min=0.0) + direction = g1 + mu * g2 + return direction * matrix[0].norm() / direction.norm() + + +@mark.parametrize(["aggregator", "matrix"], scaled_pairs + typical_pairs) +def test_expected_structure(aggregator: PCD, matrix: Tensor) -> None: + assert_expected_structure(aggregator, matrix) + + +@mark.parametrize(["aggregator", "matrix"], requires_grad_pairs) +def test_non_differentiable(aggregator: PCD, matrix: Tensor) -> None: + assert_non_differentiable(aggregator, matrix) + + +def test_representations() -> None: + A = PCD(tau=0.1, beta=0.9, eps=1e-6) + assert repr(A) == "PCD(tau=0.1, beta=0.9, eps=1e-06)" + assert str(A) == "PCD" + + W = PCDWeighting(tau=0.1, beta=0.9, eps=1e-6) + assert repr(W) == "PCDWeighting(tau=0.1, beta=0.9, eps=1e-06)" + + +def test_zero_rows_returns_zero_vector() -> None: + out = PCD()(tensor_([]).reshape(0, 3)) + assert_close(out, zeros_(3)) + + +def test_zero_columns_returns_zero_vector() -> None: + out = PCD()(tensor_([]).reshape(2, 0)) + assert out.shape == (0,) + + +def test_single_row_returns_it() -> None: + J = randn_((1, 5)) + assert_close(PCD()(J), J[0]) + + +@mark.parametrize("shape", [(2, 5), (3, 5), (5, 10), (9, 11)]) +def test_output_has_the_norm_of_the_primary_gradient(shape: tuple[int, int]) -> None: + J = randn_(shape) * randn_((shape[0], 1)).exp() + out = PCD()(J) + assert_close(out.norm(), J[0].norm()) + + +def test_zero_primary_gradient_returns_zero_vector() -> None: + J = randn_((3, 5)) + J[0] = 0.0 + assert_close(PCD()(J), zeros_(5)) + + +def test_primary_gradient_is_returned_when_constraints_are_satisfied() -> None: + J = tensor_([[1.0, 0.0, 0.0], [1.0, 1.0, 0.0], [1.0, 0.0, 1.0]]) + assert_close(PCD(tau=0.5)(J), J[0]) + + +def test_primary_gradient_is_returned_when_constraints_are_infeasible() -> None: + # The two secondary gradients are anti-parallel, so no direction improves both of them. + J = tensor_([[1.0, 1.0], [0.0, 1.0], [0.0, -2.0]]) + assert_close(PCD()(J), J[0]) + + +def test_primary_objective_is_sacrificed_for_an_opposed_secondary_objective() -> None: + # The secondary constraint can only be satisfied by moving against the primary gradient, and + # the output is then rescaled to the norm of the primary gradient. + J = tensor_([[1.0, 0.0], [-3.0, 0.0]]) + assert_close(PCD(tau=0.02)(J), tensor_([-1.0, 0.0])) + + +@mark.parametrize("tau", [0.0, 0.02, 0.5, 1.0]) +def test_two_objectives_matches_closed_form(tau: float) -> None: + J = randn_((2, 6)) + scales = 1.0 / J.norm(dim=1) + assert_close(PCD(tau=tau, eps=0.0)(J), _two_objectives_pcd(J, tau, scales)) + + +def test_second_call_uses_debiased_moving_average() -> None: + beta = 0.9 + J1 = randn_((2, 6)) + J2 = 10.0 * randn_((2, 6)) + A = PCD(tau=0.3, beta=beta, eps=0.0) + A(J1) + out = A(J2) + + ema = beta * (1 - beta) * J1.norm(dim=1) ** 2 + (1 - beta) * J2.norm(dim=1) ** 2 + scales = 1.0 / (ema / (1 - beta**2)).sqrt() + assert_close(out, _two_objectives_pcd(J2, 0.3, scales)) + + +def test_first_call_is_invariant_to_positive_row_scaling() -> None: + J = randn_((4, 6)) + c = tensor_([2.0, 0.1, 30.0, 0.5]) + assert_close(PCD(eps=0.0)(c.unsqueeze(1) * J), c[0] * PCD(eps=0.0)(J)) + + +@mark.parametrize("m", [2, 3, 5, 8]) +@mark.parametrize("tau", [0.0, 0.02, 0.5]) +def test_solve_qp_satisfies_kkt_conditions(m: int, tau: float) -> None: + """ + Tests that the weights w = [1, mu_2, ..., mu_m] satisfy the KKT conditions of the QP. Since the + QP is convex, they are sufficient for optimality. Stationarity holds by construction of w. + """ + + J = randn_((m, 10)).to(device="cpu", dtype=torch.float64) + G = J @ J.T + taus = torch.full([m - 1], tau, dtype=torch.float64) + weights = _solve_qp(G, taus) + mu = weights[1:] + slacks = G[1:] @ weights - taus * G.diagonal()[1:] + + assert weights[0] == 1.0 + assert (mu >= 0.0).all() + assert (slacks >= -1e-8).all() + assert_close(mu * slacks, torch.zeros_like(mu), rtol=0.0, atol=1e-8) + + +def test_vector_tau_with_equal_values_matches_scalar_tau() -> None: + J = randn_((4, 6)) + assert_close(PCD(tau=tensor_([0.3, 0.3, 0.3]))(J), PCD(tau=0.3)(J)) + + +def test_vector_tau_with_wrong_length_raises() -> None: + A = PCD(tau=tensor_([0.1, 0.2])) + with raises(ValueError, match="tau"): + A(randn_((4, 6))) + + +@mark.parametrize("matrix", typical_matrices_2_plus_rows) +def test_reset_restores_first_step_behavior(matrix: Tensor) -> None: + A = PCD() + first = A(matrix) + A(2.0 * matrix + 1.0) + A.reset() + assert_close(first, A(matrix)) + + +def test_weighting_reset_restores_first_step_behavior() -> None: + J = randn_((3, 8)) + G = J @ J.T + W = PCDWeighting() + first = W(G) + W(4.0 * G) + W.reset() + assert_close(first, W(G)) + + +def test_changing_m_auto_resets() -> None: + J = randn_((3, 8)) + A = PCD() + A(randn_((4, 8))) + assert_close(A(J), PCD()(J)) + + +def test_aggregator_and_weighting_agree() -> None: + A = PCD(tau=0.1) + W = PCDWeighting(tau=0.1) + for _ in range(3): + J = randn_((3, 8)) + assert_close(W(J @ J.T) @ J, A(J)) + + +def test_tau_setter_accepts_valid() -> None: + A = PCD() + A.tau = 0.0 + assert A.tau == 0.0 + A.tau = 1.0 + assert A.tau == 1.0 + tau = tensor_([0.1, 0.5]) + A.tau = tau + assert A.tau is tau + assert A.gramian_weighting.tau is tau + + +@mark.parametrize("tau", [-0.1, 1.1, float("nan")]) +def test_tau_setter_rejects_out_of_range(tau: float) -> None: + A = PCD() + with raises(ValueError, match="tau"): + A.tau = tau + with raises(ValueError, match="tau"): + A.tau = tensor_([0.1, tau]) + + +def test_tau_setter_rejects_non_vector_tensor() -> None: + A = PCD() + with raises(ValueError, match="tau"): + A.tau = tensor_([[0.1, 0.2]]) + + +def test_beta_setter_accepts_valid() -> None: + A = PCD() + A.beta = 0.0 + assert A.beta == 0.0 + A.beta = 0.5 + assert A.beta == 0.5 + assert A.gramian_weighting.beta == 0.5 + + +@mark.parametrize("beta", [-0.1, 1.0]) +def test_beta_setter_rejects_out_of_range(beta: float) -> None: + A = PCD() + with raises(ValueError, match="beta"): + A.beta = beta + + +def test_eps_setter_accepts_valid() -> None: + A = PCD() + A.eps = 0.0 + assert A.eps == 0.0 + A.eps = 1e-6 + assert A.eps == 1e-6 + assert A.gramian_weighting.eps == 1e-6 + + +def test_eps_setter_rejects_negative() -> None: + A = PCD() + with raises(ValueError, match="eps"): + A.eps = -1e-9 diff --git a/tests/unit/aggregation/test_values.py b/tests/unit/aggregation/test_values.py index a73b5d901..9c3c807bb 100644 --- a/tests/unit/aggregation/test_values.py +++ b/tests/unit/aggregation/test_values.py @@ -11,6 +11,7 @@ from torchjd.aggregation import ( IMTLG, MGDA, + PCD, Aggregator, AlignedMTL, AlignedMTLWeighting, @@ -33,6 +34,7 @@ MeanWeighting, MGDAWeighting, NashMTL, + PCDWeighting, PCGrad, PCGradWeighting, Random, @@ -74,6 +76,7 @@ (Krum(n_byzantine=1, n_selected=4), J_Krum, tensor([1.2500, 0.7500, 1.5000])), (Mean(), J_base, tensor([1.0, 1.0, 1.0])), (MGDA(), J_base, tensor([0.0, 1.0, 1.0])), + (PCD(), J_base, tensor([-0.8200, 2.9434, 2.9434])), (PCGrad(), J_base, tensor([0.5848, 3.8012, 3.8012])), (Random(), J_base, tensor([-2.6229, 1.0000, 1.0000])), (Sum(), J_base, tensor([2.0, 2.0, 2.0])), @@ -91,6 +94,7 @@ (GradVacWeighting(), G_base, tensor([2.2222, 1.5789])), (MeanWeighting(), G_base, tensor([0.5000, 0.5000])), (MGDAWeighting(), G_base, tensor([0.6000, 0.4000])), + (PCDWeighting(), G_base, tensor([1.8481, 1.0954])), (PCGradWeighting(), G_base, tensor([2.2222, 1.5789])), (RandomWeighting(), G_base, tensor([0.8623, 0.1377])), (SumWeighting(), G_base, tensor([1.0, 1.0])), From b707bd8993ecee8454624bd3271a99721fe1f8d9 Mon Sep 17 00:00:00 2001 From: dara <59355936+DaraVaram@users.noreply.github.com> Date: Tue, 29 Sep 2026 12:55:20 +0400 Subject: [PATCH 2/5] fix(aggregation): Propagate non-finite inputs in PCD and use the state key pattern - nan or inf in the Gramian now gives nan weights and leaves the moving average unchanged. Before, it returned zeros and corrupted the moving average, so every later call returned zero until reset(). - Track the state with _state_key and _ensure_state, as in GradVac. - Define the symbols of the docstring in a bullet list. Co-Authored-By: Claude Opus 5.5 --- src/torchjd/aggregation/_pcd.py | 38 ++++++++++++++++++++++-------- tests/unit/aggregation/test_pcd.py | 16 +++++++++++++ 2 files changed, 44 insertions(+), 10 deletions(-) diff --git a/src/torchjd/aggregation/_pcd.py b/src/torchjd/aggregation/_pcd.py index 5e11893d9..4e8ac1bab 100644 --- a/src/torchjd/aggregation/_pcd.py +++ b/src/torchjd/aggregation/_pcd.py @@ -46,6 +46,7 @@ def __init__(self, tau: float | Tensor = 0.02, beta: float = 0.999, eps: float = self.eps = eps self.register_buffer("_sq_norm_ema", None) self.register_buffer("_n_steps", None) + self._state_key: int | None = None @property def tau(self) -> float | Tensor: @@ -91,6 +92,7 @@ def reset(self) -> None: self._sq_norm_ema = None self._n_steps = None + self._state_key = None def forward(self, gramian: PSDMatrix, /) -> Tensor: # The problem only has size m x m, so we solve it on cpu and in float64. @@ -101,6 +103,10 @@ def forward(self, gramian: PSDMatrix, /) -> Tensor: return weights.to(device=gramian.device, dtype=gramian.dtype) taus = self._get_taus(m) + if not G.isfinite().all(): + # Let nan and inf propagate to the output, without corrupting the moving average. + return torch.full_like(gramian.diagonal(), torch.nan) + sq_norms = G.diagonal() scales = self._update_scales(sq_norms) @@ -131,11 +137,7 @@ def _update_scales(self, sq_norms: Tensor) -> Tensor: the scale by which each gradient should be multiplied to be normalized. """ - m = sq_norms.shape[0] - if self._sq_norm_ema is None or self._sq_norm_ema.shape[0] != m: - self._sq_norm_ema = torch.zeros(m, dtype=torch.float64) - self._n_steps = torch.tensor(0) - + self._ensure_state(sq_norms.shape[0]) sq_norm_ema = self._sq_norm_ema.to(sq_norms) n_steps = int(cast(Tensor, self._n_steps)) + 1 sq_norm_ema = self.beta * sq_norm_ema + (1.0 - self.beta) * sq_norms @@ -145,6 +147,14 @@ def _update_scales(self, sq_norms: Tensor) -> Tensor: debiased_sq_norm_ema = sq_norm_ema / (1.0 - self.beta**n_steps) return 1.0 / (debiased_sq_norm_ema + self.eps).sqrt() + def _ensure_state(self, m: int) -> None: + # The moving average is always kept on cpu and in float64, where the weights are computed, + # so it only depends on the number of objectives. + if self._state_key != m or self._sq_norm_ema is None: + self._sq_norm_ema = torch.zeros(m, dtype=torch.float64) + self._n_steps = torch.tensor(0) + self._state_key = m + def __repr__(self) -> str: return f"{self.__class__.__name__}(tau={self.tau!r}, beta={self.beta!r}, eps={self.eps!r})" @@ -169,11 +179,17 @@ class PCD(GramianWeightedAggregator, Stateful, _NonDifferentiable): \quad \text{subject to} \quad \tilde g_j^\top d \geq \tau \left\| \tilde g_j \right\|^2 \quad \text{for } j = 2, \dots, m, - where :math:`\tilde g_i = s_i g_i` is the normalized gradient of objective :math:`i`, with - :math:`s_i = 1 / \sqrt{\hat v_i + \epsilon}`, and :math:`\hat v_i` is the bias-corrected - exponential moving average, over the successive calls, of :math:`\|g_i\|^2`. The output is - :math:`\tilde d` rescaled to the norm of :math:`g_1`. Because of the normalization, :math:`\tau` - is a scale-free fraction. + where: + + - :math:`\tilde g_i = s_i g_i` is the normalized gradient of objective :math:`i`, + - :math:`s_i = 1 / \sqrt{\hat v_i + \epsilon}` is its scale, + - :math:`\hat v_i` is the bias-corrected exponential moving average of :math:`\|g_i\|^2` over + the successive calls, with decay :math:`\beta`. + + The output is :math:`\tilde d` rescaled to the norm of :math:`g_1`. Because of the + normalization, :math:`\tau` is a scale-free fraction: multiplying a secondary row of the input + by the same positive constant at every call leaves the output unchanged, up to the effect of + :math:`\epsilon`. Special cases: @@ -181,6 +197,8 @@ class PCD(GramianWeightedAggregator, Stateful, _NonDifferentiable): - If the constraints cannot be satisfied simultaneously (which requires :math:`m \geq 3`, e.g. two anti-parallel secondary gradients), they are dropped and the output is :math:`g_1`. - If :math:`\tilde d = 0`, the output is zero. + - If the input contains ``nan`` or ``inf``, the output is ``nan`` and the moving average is left + unchanged. :param tau: The fraction :math:`\tau \in [0, 1]` of normalized first-order progress guaranteed to each secondary objective. Either a float, shared by all secondary objectives, or a vector diff --git a/tests/unit/aggregation/test_pcd.py b/tests/unit/aggregation/test_pcd.py index 40929425d..b3a7d8203 100644 --- a/tests/unit/aggregation/test_pcd.py +++ b/tests/unit/aggregation/test_pcd.py @@ -176,6 +176,22 @@ def test_changing_m_auto_resets() -> None: assert_close(A(J), PCD()(J)) +@mark.parametrize("value", [float("nan"), float("inf"), -float("inf")]) +def test_non_finite_input_returns_nan_and_keeps_state(value: float) -> None: + J = randn_((3, 8)) + J_non_finite = J.clone() + J_non_finite[1, 2] = value + + A = PCD() + A(J) + assert A(J_non_finite).isnan().all() + + # The non-finite call must not have changed the moving average. + reference = PCD() + reference(J) + assert_close(A(J), reference(J)) + + def test_aggregator_and_weighting_agree() -> None: A = PCD(tau=0.1) W = PCDWeighting(tau=0.1) From 6d0d4f3adf962ce8738b0246cd4d95776f8c1218 Mon Sep 17 00:00:00 2001 From: dara <59355936+DaraVaram@users.noreply.github.com> Date: Tue, 29 Sep 2026 13:08:39 +0400 Subject: [PATCH 3/5] fix(aggregation): Fix PCD state typing and drop the README row - Cast the moving-average buffer after _ensure_state, as in GradVac, so that ty can type it. - Remove the README table row: the docs page does not exist on torchjd.org/stable until the next release, which is when the table gets updated. Co-Authored-By: Claude Opus 5.5 --- README.md | 1 - src/torchjd/aggregation/_pcd.py | 2 +- 2 files changed, 1 insertion(+), 2 deletions(-) diff --git a/README.md b/README.md index 9d53149fd..ff862bc4c 100644 --- a/README.md +++ b/README.md @@ -192,7 +192,6 @@ TorchJD provides many existing aggregators from the literature, listed in the fo | [MGDA](https://torchjd.org/stable/docs/aggregation/mgda#torchjd.aggregation.MGDA) | [MGDAWeighting](https://torchjd.org/stable/docs/aggregation/mgda#torchjd.aggregation.MGDAWeighting) | [Multiple-gradient descent algorithm (MGDA) for multiobjective optimization](https://comptes-rendus.academie-sciences.fr/mathematique/articles/10.1016/j.crma.2012.03.014/) | | - | [MoDoWeighting](https://torchjd.org/stable/docs/aggregation/modo/#torchjd.aggregation.MoDoWeighting) | [Three-Way Trade-Off in Multi-Objective Learning: Optimization, Generalization and Conflict-Avoidance](https://www.jmlr.org/papers/volume25/23-1287/23-1287.pdf) | | [NashMTL](https://torchjd.org/stable/docs/aggregation/nash_mtl#torchjd.aggregation.NashMTL) | - | [Multi-Task Learning as a Bargaining Game](https://arxiv.org/pdf/2202.01017) | -| [PCD](https://torchjd.org/stable/docs/aggregation/pcd#torchjd.aggregation.PCD) | [PCDWeighting](https://torchjd.org/stable/docs/aggregation/pcd#torchjd.aggregation.PCDWeighting) | [Not All Objectives Are Born Equal: Priority-Constrained Descent for Hierarchical Multi-Objective Optimization](https://openreview.net/forum?id=HT01yGHLEt) | | [PCGrad](https://torchjd.org/stable/docs/aggregation/pcgrad#torchjd.aggregation.PCGrad) | [PCGradWeighting](https://torchjd.org/stable/docs/aggregation/pcgrad#torchjd.aggregation.PCGradWeighting) | [Gradient Surgery for Multi-Task Learning](https://arxiv.org/pdf/2001.06782) | | [Random](https://torchjd.org/stable/docs/aggregation/random#torchjd.aggregation.Random) | [RandomWeighting](https://torchjd.org/stable/docs/aggregation/random#torchjd.aggregation.RandomWeighting) | [Reasonable Effectiveness of Random Weighting: A Litmus Test for Multi-Task Learning](https://arxiv.org/pdf/2111.10603) | | - | [SDMGradWeighting](https://torchjd.org/stable/docs/aggregation/sdmgrad#torchjd.aggregation.SDMGradWeighting) | [Direction-oriented Multi-objective Learning: Simple and Provable Stochastic Algorithms](https://arxiv.org/pdf/2305.18409) | diff --git a/src/torchjd/aggregation/_pcd.py b/src/torchjd/aggregation/_pcd.py index 4e8ac1bab..5eb92baef 100644 --- a/src/torchjd/aggregation/_pcd.py +++ b/src/torchjd/aggregation/_pcd.py @@ -138,7 +138,7 @@ def _update_scales(self, sq_norms: Tensor) -> Tensor: """ self._ensure_state(sq_norms.shape[0]) - sq_norm_ema = self._sq_norm_ema.to(sq_norms) + sq_norm_ema = cast(Tensor, self._sq_norm_ema).to(sq_norms) n_steps = int(cast(Tensor, self._n_steps)) + 1 sq_norm_ema = self.beta * sq_norm_ema + (1.0 - self.beta) * sq_norms self._sq_norm_ema = sq_norm_ema From 76e2a61544a5ac0a9b68dde52df01b5518211307 Mon Sep 17 00:00:00 2001 From: dara <59355936+DaraVaram@users.noreply.github.com> Date: Tue, 29 Sep 2026 16:10:45 +0400 Subject: [PATCH 4/5] fix(aggregation): Address review of PCD - Cap the iterations of the active-set solver at 10 m (it needs about m in practice, and the cap never binds on 53,640 random and ill-conditioned problems with up to 100 objectives). When it is reached, the current iterate is returned. - Raise the tolerance of the linear-dependence test to the precision of the Gramian (10 machine epsilons of its dtype, at least 1e-9). With float32 Gramians, rounding errors made dependent gradients look independent, which gave huge multipliers and a wrong direction. - Return zero when the direction is zero up to the rounding errors of the Gramian, instead of rescaling rounding noise to the norm of the primary gradient. - Give a zero scale to objectives whose gradient has always been zero, which avoids 1 / 0 when eps = 0. - Parametrize the norm and closed-form tests on typical and scaled matrices. - Add PCD to the interactive plotter. Co-Authored-By: Claude Opus 5.5 --- src/torchjd/aggregation/_pcd.py | 62 +++++++++++++++++++++++------ tests/plots/interactive_plotter.py | 2 + tests/unit/aggregation/test_pcd.py | 63 ++++++++++++++++++++++++------ 3 files changed, 101 insertions(+), 26 deletions(-) diff --git a/src/torchjd/aggregation/_pcd.py b/src/torchjd/aggregation/_pcd.py index 5eb92baef..9c8657f4b 100644 --- a/src/torchjd/aggregation/_pcd.py +++ b/src/torchjd/aggregation/_pcd.py @@ -110,14 +110,22 @@ def forward(self, gramian: PSDMatrix, /) -> Tensor: sq_norms = G.diagonal() scales = self._update_scales(sq_norms) + # Relative precision of the Gramian. Below it, the solver treats gradients as linearly + # dependent, and the direction as zero. + precision = 10.0 * torch.finfo(gramian.dtype).eps + # If the gradient of the primary objective is zero, the update is zero. if sq_norms[0] > 0.0: normalized_G = G * torch.outer(scales, scales) - coefficients = _solve_qp(normalized_G, taus) * scales - direction_sq_norm = coefficients @ G @ coefficients - # Rescale the direction to the norm of the primary gradient. - if direction_sq_norm > 0.0: - weights = coefficients * (sq_norms[0] / direction_sq_norm).sqrt() + w = _solve_qp(normalized_G, taus, dependence_tol=max(1e-9, precision)) + # Squared norm of d~ = sum_i w_i s_i g_i, and the scale of its rounding error when it is + # computed from the Gramian. + direction_sq_norm = w @ normalized_G @ w + magnitude = (w.abs() @ normalized_G.diagonal().sqrt()) ** 2 + # Rescale the direction to the norm of the primary gradient, unless it is zero up to + # rounding errors. + if direction_sq_norm > precision * magnitude: + weights = w * scales * (sq_norms[0] / direction_sq_norm).sqrt() return weights.to(device=gramian.device, dtype=gramian.dtype) @@ -145,7 +153,10 @@ def _update_scales(self, sq_norms: Tensor) -> Tensor: self._n_steps = torch.tensor(n_steps) debiased_sq_norm_ema = sq_norm_ema / (1.0 - self.beta**n_steps) - return 1.0 / (debiased_sq_norm_ema + self.eps).sqrt() + denominators = (debiased_sq_norm_ema + self.eps).sqrt() + # A zero moving average means that the gradient has always been zero, so its scale does not + # matter. Using 0 avoids dividing by zero when eps = 0. + return torch.where(denominators > 0.0, 1.0 / denominators, 0.0) def _ensure_state(self, m: int) -> None: # The moving average is always kept on cpu and in float64, where the weights are computed, @@ -196,7 +207,7 @@ class PCD(GramianWeightedAggregator, Stateful, _NonDifferentiable): - If :math:`g_1 = 0`, the output is zero. - If the constraints cannot be satisfied simultaneously (which requires :math:`m \geq 3`, e.g. two anti-parallel secondary gradients), they are dropped and the output is :math:`g_1`. - - If :math:`\tilde d = 0`, the output is zero. + - If :math:`\tilde d = 0`, up to the rounding errors of the Gramian, the output is zero. - If the input contains ``nan`` or ``inf``, the output is ``nan`` and the moving average is left unchanged. @@ -261,7 +272,13 @@ def __repr__(self) -> str: return f"{self.__class__.__name__}(tau={self.tau!r}, beta={self.beta!r}, eps={self.eps!r})" -def _solve_qp(G: Tensor, taus: Tensor, tol: float = 1e-9) -> Tensor: +def _solve_qp( + G: Tensor, + taus: Tensor, + tol: float = 1e-9, + dependence_tol: float = 1e-9, + max_iters: int | None = None, +) -> Tensor: r""" Solves the quadratic program of PCD, given the Gramian :math:`G` of the (normalized) gradients :math:`g_1, \dots, g_m`, whose first row and column correspond to the primary objective: @@ -278,6 +295,18 @@ def _solve_qp(G: Tensor, taus: Tensor, tol: float = 1e-9) -> Tensor: This is the dual active-set method of Goldfarb and Idnani (1983), for an identity Hessian. The stationarity condition :math:`d = g_1 + \sum_j \mu_j g_j` holds at every iteration, so that all the required inner products can be read from :math:`G`. + + A constraint counts as satisfied when it is violated by at most ``tol`` times the largest + diagonal entry of :math:`G`. A gradient counts as a linear combination of the gradients of the + active constraints when its component orthogonal to them has a squared norm of at most + ``dependence_tol`` times its own. The latter should be above the relative precision of + :math:`G`, or rounding errors can make dependent gradients look independent, with huge + multipliers. + + Each iteration adds or drops one constraint. In exact arithmetic, the method terminates after + finitely many iterations (about :math:`m` in practice). ``max_iters`` (by default :math:`10 m`) + only guards against cycling caused by rounding errors: when it is reached, the current iterate + is returned. """ m = G.shape[0] @@ -286,23 +315,27 @@ def _solve_qp(G: Tensor, taus: Tensor, tol: float = 1e-9) -> Tensor: if m == 1: return weights + if max_iters is None: + max_iters = 10 * m Gs = G[1:, 1:] b = taus * Gs.diagonal() - G[1:, 0] # b_j > 0 iff g_1 violates the constraint of objective j mu = torch.zeros(m - 1, dtype=G.dtype) atol = tol * max(G.diagonal().max().item(), torch.finfo(G.dtype).tiny) active: list[int] = [] + n_iters = 0 - while True: + while n_iters < max_iters: slacks = Gs @ mu - b slacks[active] = 0.0 # The active constraints hold with equality, up to rounding errors. p = int(slacks.argmin()) if slacks[p] >= -atol: - weights[1:] = mu.clamp(min=0.0) - return weights + break # Add the violated constraint p to the active set, possibly dropping some others first. slack_p = slacks[p] - while True: + while n_iters < max_iters: + n_iters += 1 + # In the dual space, the step direction is +1 for mu_p and -r for the active mu_j. # In the primal space, it is z = g_p - sum_j r_j g_j, the component of g_p orthogonal to # the gradients of the active constraints. @@ -320,7 +353,7 @@ def _solve_qp(G: Tensor, taus: Tensor, tol: float = 1e-9) -> Tensor: # If z = 0, g_p is a linear combination of the active gradients and no primal step is # possible. - is_z_non_zero = z_sq_norm > tol * Gs[p, p] + is_z_non_zero = z_sq_norm > dependence_tol * Gs[p, p] full_step = (-slack_p / z_sq_norm).item() if is_z_non_zero else torch.inf step = min(partial_step, full_step) @@ -337,3 +370,6 @@ def _solve_qp(G: Tensor, taus: Tensor, tol: float = 1e-9) -> Tensor: mu[active[k]] = 0.0 del active[k] slack_p = Gs[p] @ mu - b[p] + + weights[1:] = mu.clamp(min=0.0) + return weights diff --git a/tests/plots/interactive_plotter.py b/tests/plots/interactive_plotter.py index 8c15ee40e..cbcece703 100644 --- a/tests/plots/interactive_plotter.py +++ b/tests/plots/interactive_plotter.py @@ -15,6 +15,7 @@ from torchjd.aggregation import ( IMTLG, MGDA, + PCD, Aggregator, AlignedMTL, CAGrad, @@ -71,6 +72,7 @@ def main() -> None: str(Mean()): lambda: Mean(), str(MGDA()): lambda: MGDA(), str(NashMTL(n_tasks=n_tasks)): lambda: NashMTL(n_tasks=n_tasks), + str(PCD()): lambda: PCD(), str(PCGrad()): lambda: PCGrad(), str(Random()): lambda: Random(), str(Sum()): lambda: Sum(), diff --git a/tests/unit/aggregation/test_pcd.py b/tests/unit/aggregation/test_pcd.py index b3a7d8203..a3051b288 100644 --- a/tests/unit/aggregation/test_pcd.py +++ b/tests/unit/aggregation/test_pcd.py @@ -8,7 +8,12 @@ from torchjd.aggregation._pcd import _solve_qp from ._asserts import assert_expected_structure, assert_non_differentiable -from ._inputs import scaled_matrices, typical_matrices, typical_matrices_2_plus_rows +from ._inputs import ( + scaled_matrices, + scaled_matrices_2_plus_rows, + typical_matrices, + typical_matrices_2_plus_rows, +) scaled_pairs = [(PCD(), m) for m in scaled_matrices] typical_pairs = [(PCD(), m) for m in typical_matrices] @@ -21,10 +26,16 @@ def _two_objectives_pcd(matrix: Tensor, tau: float, scales: Tensor) -> Tensor: derive expected values independently of the implementation. """ - g1, g2 = matrix[0] * scales[0], matrix[1] * scales[1] - mu = ((tau * (g2 @ g2) - g1 @ g2) / (g2 @ g2)).clamp(min=0.0) + matrix64, scales64 = matrix.to(dtype=torch.float64), scales.to(dtype=torch.float64) + g1, g2 = matrix64[0] * scales64[0], matrix64[1] * scales64[1] + sq_norm_2 = g2 @ g2 + mu = ((tau * sq_norm_2 - g1 @ g2) / sq_norm_2).clamp(min=0.0) if sq_norm_2 > 0.0 else 0.0 direction = g1 + mu * g2 - return direction * matrix[0].norm() / direction.norm() + # The direction vanishes when the primary gradient is zero, or when tau = 0 and the secondary + # gradient is opposite to it. Up to rounding errors, it is then zero. + if direction.norm() <= 1e-6 * (g1.norm() + mu * g2.norm()): + return torch.zeros_like(matrix[0]) + return (direction * matrix64[0].norm() / direction.norm()).to(dtype=matrix.dtype) @mark.parametrize(["aggregator", "matrix"], scaled_pairs + typical_pairs) @@ -61,11 +72,10 @@ def test_single_row_returns_it() -> None: assert_close(PCD()(J), J[0]) -@mark.parametrize("shape", [(2, 5), (3, 5), (5, 10), (9, 11)]) -def test_output_has_the_norm_of_the_primary_gradient(shape: tuple[int, int]) -> None: - J = randn_(shape) * randn_((shape[0], 1)).exp() - out = PCD()(J) - assert_close(out.norm(), J[0].norm()) +@mark.parametrize("matrix", typical_matrices + scaled_matrices) +def test_output_has_the_norm_of_the_primary_gradient(matrix: Tensor) -> None: + out = PCD()(matrix) + assert_close(out.norm(), matrix[0].norm(), rtol=1e-4, atol=0.0) def test_zero_primary_gradient_returns_zero_vector() -> None: @@ -92,11 +102,17 @@ def test_primary_objective_is_sacrificed_for_an_opposed_secondary_objective() -> assert_close(PCD(tau=0.02)(J), tensor_([-1.0, 0.0])) +@mark.parametrize("matrix", typical_matrices_2_plus_rows + scaled_matrices_2_plus_rows) @mark.parametrize("tau", [0.0, 0.02, 0.5, 1.0]) -def test_two_objectives_matches_closed_form(tau: float) -> None: - J = randn_((2, 6)) - scales = 1.0 / J.norm(dim=1) - assert_close(PCD(tau=tau, eps=0.0)(J), _two_objectives_pcd(J, tau, scales)) +def test_two_objectives_matches_closed_form(matrix: Tensor, tau: float) -> None: + J = matrix[:2] + eps = 1e-8 + scales = 1.0 / (J.norm(dim=1) ** 2 + eps).sqrt() + # The tolerance is relative to the norm of the output, so that it also applies to scaled matrices. + atol = 1e-4 * max(1.0, J[0].norm().item()) + assert_close( + PCD(tau=tau, eps=eps)(J), _two_objectives_pcd(J, tau, scales), rtol=1e-4, atol=atol + ) def test_second_call_uses_debiased_moving_average() -> None: @@ -139,6 +155,27 @@ def test_solve_qp_satisfies_kkt_conditions(m: int, tau: float) -> None: assert_close(mu * slacks, torch.zeros_like(mu), rtol=0.0, atol=1e-8) +def test_solve_qp_stops_after_max_iters() -> None: + # g_1 violates both constraints, and adding the first one to the active set is not enough, so the + # solver needs two iterations. + J = tensor_([[1.0, 0.0], [-1.0, 1.0], [-1.0, -1.0]]).to(device="cpu", dtype=torch.float64) + G = J @ J.T + taus = torch.full([2], 0.5, dtype=torch.float64) + + weights = _solve_qp(G, taus) + assert (weights[1:] > 0.0).all() + + truncated_weights = _solve_qp(G, taus, max_iters=1) + assert truncated_weights[0] == 1.0 + assert (truncated_weights[1:] > 0.0).sum() == 1 + + +def test_zero_secondary_gradient_with_zero_eps() -> None: + # The moving average of the second row is zero, so its scale must not become 1 / 0. + J = tensor_([[1.0, 2.0], [0.0, 0.0], [3.0, -1.0]]) + assert_close(PCD(eps=0.0)(J), J[0]) + + def test_vector_tau_with_equal_values_matches_scalar_tau() -> None: J = randn_((4, 6)) assert_close(PCD(tau=tensor_([0.3, 0.3, 0.3]))(J), PCD(tau=0.3)(J)) From f284a2cc46301c6482af20766387e9ad448f7bc6 Mon Sep 17 00:00:00 2001 From: dara <59355936+DaraVaram@users.noreply.github.com> Date: Tue, 29 Sep 2026 18:10:10 +0400 Subject: [PATCH 5/5] Update tests/unit/aggregation/test_pcd.py MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: Valérian Rey <31951177+ValerianRey@users.noreply.github.com> --- tests/unit/aggregation/test_pcd.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/unit/aggregation/test_pcd.py b/tests/unit/aggregation/test_pcd.py index a3051b288..2705be654 100644 --- a/tests/unit/aggregation/test_pcd.py +++ b/tests/unit/aggregation/test_pcd.py @@ -75,7 +75,7 @@ def test_single_row_returns_it() -> None: @mark.parametrize("matrix", typical_matrices + scaled_matrices) def test_output_has_the_norm_of_the_primary_gradient(matrix: Tensor) -> None: out = PCD()(matrix) - assert_close(out.norm(), matrix[0].norm(), rtol=1e-4, atol=0.0) + assert_close(out.norm(), matrix[0].norm(), rtol=2e-4, atol=0.0) def test_zero_primary_gradient_returns_zero_vector() -> None: