function

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