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
| Name | Type | Description |
|---|---|---|
| P | torch.Tensor | Materialized transport plan (n, m). |
| diag_x | torch.Tensor | Diagonal scaling for source marginal (n,). |
| denom | torch.Tensor | Diagonal of the Schur complement (m,), includes regularization. |
| rhs | torch.Tensor | Right-hand side vector (m,). |
| max_iter = 100 | int | Maximum CG iterations. |
| rtol = 1e-06 | float | Relative tolerance. |
| atol = 1e-06 | float | Absolute tolerance. |
Returns
z — Solution vector (m,).
Source: torchmatch/transport/samples/kernels/cg_dense.py:91