waxmorph.jax.losses

waxmorph.jax.losses#

JAX losses for ordered arrays and unordered point clouds.

squared_loss uses row correspondence; Chamfer and Sinkhorn support unordered clouds. The OTT factory provides debiased Sinkhorn divergence; PyTorch also provides GeomLoss families.

Functions

chamfer_distance(X_pred, X_target)

Two directional mean distances between unordered, possibly unequal clouds.

make_sinkhorn_loss([blur, p, cost_fn])

Build OTT's debiased entropic-OT divergence for unordered clouds.

squared_loss(X_pred, X_target)

Squared Frobenius norm for row-aligned arrays of equal shape.