function

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

NameTypeDescription
xtorch.TensorSource points (n, d).
ytorch.TensorTarget points (m, d).
ftorch.TensorSource potential (n,).
gtorch.TensorTarget potential (m,).
diag_xtorch.TensorDiagonal D_x (n,) — row sums of P.
denomtorch.TensorDenominator D_y (m,) — column sums of P plus regularization.
rhstorch.TensorRight-hand side b (m,).
epsfloatRegularization parameter.
x2 = Nonetorch.Tensor | NonePrecomputed ``||x||²`` (n,); computed on the fly if None.
y2 = Nonetorch.Tensor | NonePrecomputed ``||y||²`` (m,); computed on the fly if None.
x0 = Nonetorch.Tensor | NoneInitial CG guess (m,); zero-initialized if None.
max_iter = 50intMaximum CG iterations.
rtol = 1e-06floatRelative convergence tolerance.
atol = 1e-06floatAbsolute convergence tolerance.
autotune = FalseboolWhether to enable Triton autotuning for apply_plan_vec.

Returns

sol — Solution vector (m,).

Source: torchmatch/transport/samples/kernels/cg_python_batched.py:38