Skip to content

[MRG] Fix GMM rand map overflow - #872

Open
jonathan-legrand wants to merge 14 commits into
PythonOT:masterfrom
jonathan-legrand:gmm-map-overflow
Open

jonathan-legrand wants to merge 14 commits into
PythonOT:masterfrom
jonathan-legrand:gmm-map-overflow

Conversation

@jonathan-legrand

Copy link
Copy Markdown

Types of changes

I simplified the computation of the $T_{rand}$ map and used the logsumexp trick to avoid overflows. I also added a test case that triggers overflow errors with the current main branch implementation and passes with the changes introduced in this request.

Motivation and context / Related issue

I ran into overflow errors when transporting gaussian mixtures with the $T_{rand}$ map. They occur when a point of the source domain $x$ is far away from one of the target components.
The current implementation computes: log_diff = log_g[:, None] - log_g[None, :] which is effectively:

$$ \log( g_i(x) ) - \log(g_j(x)) = \log( \frac{g_i(x)} {g_j(x) }) $$

and then exponentiates this quantity : weighted_exp = w_s[:, None] * nx.exp(log_diff)

When the ratio $g_i(x)/g_j(x)$ is very large, the exp overflows.

I rewrote this computation using the logsumexp trick. I saw that a logsumexp function is available through nx but it does not accept weights, which are required for computing the $T_{rand}$ attribution probability, so I implemented one which is backend agnostic.

How has this been tested (if it applies)

The logsumexp has been tested for high (710) and low (0) logits and it behaves as expected. The tests pass for numpy and torch backends. I had to skip jax tests because the functions of the gmm module use a lot of array assignments.

PR checklist

  • I have read the CONTRIBUTING document.
  • The documentation is up-to-date with the changes I made (check build artifacts).
  • All tests passed, and additional code has been covered with new tests.
  • I have added the PR and Issue fix to the RELEASES.md file.

@codecov

codecov Bot commented Oct 3, 2026 •

Copy link
Copy Markdown

Codecov Report

❌ Patch coverage is 98.00000% with 1 line in your changes missing coverage. Please review.
✅ Project coverage is 97.00%. Comparing base (604c47f) to head (12bb757).

Additional details and impacted files
@@           Coverage Diff           @@
##           master     #872   +/-   ##
=======================================
  Coverage   97.00%   97.00%           
=======================================
  Files         128      128           
  Lines       26349    26390   +41     
=======================================
+ Hits        25559    25599   +40     
- Misses        790      791    +1     
🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

@rflamary
rflamary requested a review from eloitanguy October 5, 2026 11:58

@eloitanguy eloitanguy left a comment •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Thank you very much for this useful PR! I only have minor comments, this is very well done (I'm impressed with the rigorous testing).

Don't forget to update RELEASES.md!

Comment thread ot/gmm.py Outdated
Comment thread ot/gmm.py Outdated
@jonathan-legrand

Copy link
Copy Markdown
Author

Thanks for your kind message! I renamed the denom and I moved its computation out of the loop. It does not even depend on i actually, I should have moved it from the very start.

The dictionary idea could avoid a few Cs12 inversions though, but as you said the gain is not obvious for most cases so maybe it's better to keep the code simpler

Comment thread ot/gmm.py
return emd(w_s, w_t, D, log=log)


def logsumexp(a, scaling_factor, axis=None):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Hello instead of rewriting a logsumexp here it woudl be better to use the one form the backend nx.logsumext it probably does not have scaling_factor but that is an easy update (in backend.py" no?

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

agreed!

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants