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/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..9c8657f4b --- /dev/null +++ b/src/torchjd/aggregation/_pcd.py @@ -0,0 +1,375 @@ +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) + self._state_key: int | None = 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 + 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. + 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) + 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) + + # 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) + 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) + + 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. + """ + + self._ensure_state(sq_norms.shape[0]) + 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 + self._n_steps = torch.tensor(n_steps) + + debiased_sq_norm_ema = sq_norm_ema / (1.0 - self.beta**n_steps) + 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, + # 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})" + + +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`, + - :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: + + - 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`, 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. + + :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, + 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: + + .. 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`. + + 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] + weights = torch.zeros(m, dtype=G.dtype) + weights[0] = 1.0 + 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 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: + break + + # Add the violated constraint p to the active set, possibly dropping some others first. + slack_p = slacks[p] + 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. + 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 > 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) + 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] + + 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 new file mode 100644 index 000000000..2705be654 --- /dev/null +++ b/tests/unit/aggregation/test_pcd.py @@ -0,0 +1,295 @@ +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, + 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] +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. + """ + + 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 + # 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) +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("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=2e-4, atol=0.0) + + +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("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(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: + 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_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)) + + +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)) + + +@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) + 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])),