function

transport.samples.kernels.cg_dense.materialize_transport_plan

def materialize_transport_plan(x: torch.Tensor, y: torch.Tensor, f_hat: torch.Tensor, g_hat: torch.Tensor, eps: float, x2: torch.Tensor | None = None, y2: torch.Tensor | None = None) -> torch.Tensor

Materialize the transport plan P = exp((f_hat + g_hat - C) / eps).

Parameters

NameTypeDescription
xtorch.TensorSource points (n, d)
ytorch.TensorTarget points (m, d)
f_hattorch.TensorSource raw-form potential (n,)
g_hattorch.TensorTarget raw-form potential (m,)
epsfloatEntropy regularization
x2 = Nonetorch.Tensor | NonePrecomputed ||x_i||² (n,), optional
y2 = Nonetorch.Tensor | NonePrecomputed ||y_j||² (m,), optional

Returns

torch.Tensor — Transport plan matrix (n, m), dtype=float32

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