function

transport.samples.kernels.cg_dense.dense_cg_solve

def dense_cg_solve(P: torch.Tensor, diag_x: torch.Tensor, denom: torch.Tensor, rhs: torch.Tensor, max_iter: int = 100, rtol: float = 1e-06, atol: float = 1e-06) -> tuple[torch.Tensor, DenseCgInfo]

Solve the Schur complement system using dense matrix operations.

Solves: (denom * I - P.T @ diag(1/diag_x) @ P) @ z = rhs This is the inner CG solve for the HVP. Instead of using streaming Triton kernels to apply the transport plan, we cache P and use dense matvecs.

Parameters

NameTypeDescription
Ptorch.TensorMaterialized transport plan (n, m).
diag_xtorch.TensorDiagonal scaling for source marginal (n,).
denomtorch.TensorDiagonal of the Schur complement (m,), includes regularization.
rhstorch.TensorRight-hand side vector (m,).
max_iter = 100intMaximum CG iterations.
rtol = 1e-06floatRelative tolerance.
atol = 1e-06floatAbsolute tolerance.

Returns

z — Solution vector (m,).

Source: torchmatch/transport/samples/kernels/cg_dense.py:91