Embedders¶
- class stix.core.embedder.Embedder(*args, **kwargs)¶
-
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:
- 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:
- class stix.core.embedder.IdentityEmbedder(*args, **kwargs)¶
Bases:
EmbedderIdentity 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:
- from_embeddings_to_raw(embedded_var)¶
Return embedded data unchanged.
- Parameters:
embedded_var (EmbeddedVar) – The embedded variable.
- Returns:
embedded_var, unchanged.- Return type:
- class stix.core.embedder.OneHotDiscreteEmbedder(*args, **kwargs)¶
Bases:
EmbedderParameter-free one-hot embedder for discrete inputs.
Maps integer indices to one-hot vectors, and passes existing one-hot inputs through unchanged. Decoding:
argmaxover the last axis. No learnable parameters. There are two construction modes, distinguished by whethernum_statesis 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 equalsdm_shape.Index-valued raw data (
num_statesgiven): the raw data are integer indices of shapedm_shape(e.g. a length-Ltoken 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, settingnum_states = K + 1so the mask symbol (indexK) 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:
- 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:
- class stix.core.embedder.LearnedDiscreteEmbedder(*args, **kwargs)¶
Bases:
EmbedderLearnable embedding table mapping discrete inputs to continuous embeddings.
Decoding from embeddings to raw discrete indices proceeds as follows:
The dot product of the embedding table and the embedded variable is computed and interpreted as logits.
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
OneHotDiscreteEmbedderunless 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:
- 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:
- 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_varwith the embedding table.- Return type:
- 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: