transport.samples.kernels.cg_python_batched.python_batched_cg_solve
def python_batched_cg_solve(x: torch.Tensor, y: torch.Tensor, f: torch.Tensor, g: torch.Tensor, diag_x: torch.Tensor, denom: torch.Tensor, rhs: torch.Tensor, eps: float, x2: torch.Tensor | None = None, y2: torch.Tensor | None = None, x0: torch.Tensor | None = None, max_iter: int = 50, rtol: float = 1e-06, atol: float = 1e-06, autotune: bool = False) -> tuple[torch.Tensor, PythonBatchedCgInfo]Solve H @ x = b using CG with external apply_plan_vec kernels.
This is a simpler, more reliable implementation that uses the proven
apply_plan_vec_sqeuclid kernels instead of inline Triton softmax.
The linear operator is:
H @ v = denom * v - P^T @ (P @ v / diag_x)
where P is the (n, m) transport plan matrix.
Parameters
| Name | Type | Description |
|---|---|---|
| x | torch.Tensor | Source points (n, d). |
| y | torch.Tensor | Target points (m, d). |
| f | torch.Tensor | Source potential (n,). |
| g | torch.Tensor | Target potential (m,). |
| diag_x | torch.Tensor | Diagonal D_x (n,) — row sums of P. |
| denom | torch.Tensor | Denominator D_y (m,) — column sums of P plus regularization. |
| rhs | torch.Tensor | Right-hand side b (m,). |
| eps | float | Regularization parameter. |
| x2 = None | torch.Tensor | None | Precomputed ``||x||²`` (n,); computed on the fly if None. |
| y2 = None | torch.Tensor | None | Precomputed ``||y||²`` (m,); computed on the fly if None. |
| x0 = None | torch.Tensor | None | Initial CG guess (m,); zero-initialized if None. |
| max_iter = 50 | int | Maximum CG iterations. |
| rtol = 1e-06 | float | Relative convergence tolerance. |
| atol = 1e-06 | float | Absolute convergence tolerance. |
| autotune = False | bool | Whether to enable Triton autotuning for apply_plan_vec. |
Returns
sol — Solution vector (m,).
Source: torchmatch/transport/samples/kernels/cg_python_batched.py:38