Coupling

class stix.core.coupling.Coupling

Bases: ABC

Abstract base class for the coupling function that correlates the source and target variables of a single batch.

Couplings allow to simulate a sampling of the source and target variables from a joint distribution \(\pi(x_{\mathrm{src}}, x_{\mathrm{tgt}})\) (or, equivalently, \(\pi(z_{\mathrm{src}}, z_{\mathrm{tgt}})\)).

This can only be used for online coupling, that is, couplings that operate on each batch independently. Offline couplings that require to access a full dataset at once must be implemented using custom data loaders.

abstractmethod __call__(raw_pairs, embedded_pairs)

Couple the raw and embedded pairs.

The same permutation (or other coupling function) is applied to the raw and embedded pairs, so both stay aligned.

Parameters:
Returns:

The re-coupled (raw_pairs, embedded_pairs).

Return type:

tuple[PyTree[stix.typing.variables.RawSourceTargetPair], PyTree[stix.typing.variables.EmbeddedSourceTargetPair]]