transport.samples.kernels.streaming_sqeuclid.standard_to_shifted_potentials
def standard_to_shifted_potentials(f: torch.Tensor, g: torch.Tensor, alpha: torch.Tensor, beta: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]Convert standard potentials to shifted form.
Args:
f: Standard source potential [n]
g: Standard target potential [m]
alpha: Source squared norms [n] = cost_scale * ||x||²
beta: Target squared norms [m] = cost_scale * ||y||²
Returns:
f_hat: Shifted source potential [n] = f - alpha
g_hat: Shifted target potential [m] = g - beta
Source: torchmatch/transport/samples/kernels/streaming_sqeuclid.py:1532