Embedders

class stix.core.embedder.Embedder(*args, **kwargs)

Bases: Module, ABC

Base class for embedder and de-embedder maps between raw and embedded source/target pairs.

Maps the raw data source/target pairs \((x_{\mathrm{src}}, x_{\mathrm{tgt}})\) to the embedded source/target pairs \((z_{\mathrm{src}}, z_{\mathrm{tgt}})\) and vice versa.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

abstractmethod from_raw_to_embeddings(raw_var)

Encode raw data into the embedding space.

Parameters:

raw_var (RawVar) – The raw variable to encode.

Returns:

The encoded variable in embedding space.

Return type:

EmbeddedVar

abstractmethod from_embeddings_to_raw(embedded_var)

Decode embeddings back to raw data space.

Parameters:

embedded_var (EmbeddedVar) – The embedded variable to decode.

Returns:

The decoded variable in raw space.

Return type:

RawVar

class stix.core.embedder.IdentityEmbedder(*args, **kwargs)

Bases: Embedder

Identity embedder: a no-op that returns the input as the embedded variable.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

from_raw_to_embeddings(raw_var)

Return raw data unchanged.

Parameters:

raw_var (RawVar) – The raw variable.

Returns:

raw_var, unchanged.

Return type:

EmbeddedVar

from_embeddings_to_raw(embedded_var)

Return embedded data unchanged.

Parameters:

embedded_var (EmbeddedVar) – The embedded variable.

Returns:

embedded_var, unchanged.

Return type:

RawVar

class stix.core.embedder.OneHotDiscreteEmbedder(*args, **kwargs)

Bases: Embedder

Parameter-free one-hot embedder for discrete inputs.

Maps integer indices to one-hot vectors, and passes existing one-hot inputs through unchanged. Decoding: argmax over the last axis. No learnable parameters. There are two construction modes, distinguished by whether num_states is passed explicitly:

  • One-hot raw data (num_states=None, the default): dm_shape’s last axis is the category count, so the raw data are already one-hot vectors (or a single categorical index that one-hots to that width) and the embedded shape equals dm_shape.

  • Index-valued raw data (num_states given): the raw data are integer indices of shape dm_shape (e.g. a length-L token sequence, dm_shape=(L,), or a scalar index, dm_shape=()), and the one-hot axis is appended, so the embedded shape is (*dm_shape, num_states). This mode can be used, for instance, for masked diffusion, setting num_states = K + 1 so the mask symbol (index K) gets its own slot.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

from_raw_to_embeddings(raw_var)

Encode discrete indices to one-hot vectors.

Parameters:

raw_var (RawVar) – Integer indices, or already-one-hot vectors that are passed through unchanged.

Returns:

One-hot vectors of width num_states.

Return type:

EmbeddedVar

from_embeddings_to_raw(embedded_var)

Decode to an integer index using an argmax.

Parameters:

embedded_var (EmbeddedVar) – The one-hot embedded variable. We note that predictions may be passed in here which may not strictly be one-hot.

Returns:

The argmax category index over the last axis.

Return type:

RawDisVar

class stix.core.embedder.LearnedDiscreteEmbedder(*args, **kwargs)

Bases: Embedder

Learnable embedding table mapping discrete inputs to continuous embeddings.

Decoding from embeddings to raw discrete indices proceeds as follows:

  1. The dot product of the embedding table and the embedded variable is computed and interpreted as logits.

  2. Decoding to an integer index is done using an argmax on the probabilities corresponding to the logits.

Warning

This embedder is experimental and potentially unstable. Learned discrete embeddings require careful stabilisation (e.g. noise schedules, codebook normalisation) that has not been implemented here — see CDCD (Dieleman et al., 2022) for discussion. Prefer OneHotDiscreteEmbedder unless you have a specific reason to use learned embeddings.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

from_raw_to_embeddings(raw_var)

Encode discrete indices (or one-hot vectors) into continuous embeddings.

Parameters:

raw_var (RawVar) – Integer indices, or one-hot vectors that are first argmaxed.

Returns:

The looked-up embeddings.

Return type:

EmbeddedVar

from_embeddings_to_raw(embedded_var)

Decode embeddings to an integer index.

Parameters:

embedded_var (EmbeddedVar) – The embedded variable to decode.

Returns:

The argmax category index.

Return type:

RawDisVar

from_embeddings_to_logits(embedded_var)

Decode embeddings to logits over the categories.

Parameters:

embedded_var (EmbeddedVar) – The embedded variable to decode.

Returns:

Logits, the dot product of embedded_var with the embedding table.

Return type:

Var

from_embeddings_to_probs(embedded_var)

Decode embeddings to a probability distribution over the categories.

Parameters:

embedded_var (EmbeddedVar) – The embedded variable to decode.

Returns:

Softmax probabilities over the categories.

Return type:

Var