function

transport.matrix.ops.sinkhorn_divergence

def sinkhorn_divergence(cost: torch.Tensor, eps: float, n_iter: int, a: torch.Tensor, b: torch.Tensor, mask: torch.Tensor | None = None, scaling: float | None = None, cost_aa: torch.Tensor | None = None, cost_bb: torch.Tensor | None = None) -> torch.Tensor

Source: torchmatch/transport/matrix/ops.py:18