Embedded source prior

class stix.core.embedded_source_prior.EmbeddedSourcePrior(*args, **kwargs)

Bases: Protocol

Protocol for sampling an embedded source \(z_{\mathrm{src}}\) from a prior.

For modalities with a two-sided interpolant, this is can be used to fill missing embedded_pairs.source values at training and sampling time.

NOTE: runtime_checkable so a prior can be recognised by isinstance

alongside the other modality fields. The check is structural: it tests that __call__ exists, never its signature.

__call__(key, shape)

Sample an embedded source \(z_{\mathrm{src}}\) of the given shape.

Parameters:
  • key (Array) – PRNG key for the draw.

  • shape (Shape) – Shape of the embedded source to draw (typically (batch, *embedding_shape)).

Returns:

An embedded source sample \(z_{\mathrm{src}}\).

Return type:

EmbeddedVar