Context encoders

class stix.nn.SumContextEncoder(*args, **kwargs)

Bases: Module

Sum per-source context vectors into a single context.

Pipeline: time_encoder(t) plus, for each encoder in context_encoders, its jax.tree.map over the matching leaf of context_data, gated by the matching leaf of context_mask; all terms are then summed. context_encoders is an arbitrary pytree of modules and context_data / context_mask must share its structure, with the encoders as its leaves. All encoder outputs must share the same feature dim so the sum is well-defined.

Conditioning semantics:

  • Whole-tree: pass context_data=None — every conditioning encoder is skipped, context is the time term alone. Cleanest path, used internally by get_classifier_free_guidance_generator() for its unconditional branch.

  • Per-leaf zero: a context_mask leaf of 0 with any (real-valued) context_data leaf — the encoder runs but its contribution is multiplied by zero. Useful when some sources are conditional and others not in the same forward.

  • Per-leaf None: a context_mask leaf of None (or whole None) means fully conditional for that source.

Parameters:
  • args (Any)

  • kwargs (Any)

Return type:

Any

__call__(t, context_data=None, context_mask=None)

Build the context vector.

Parameters:
  • t (Time) – Scalar time fed to time_encoder.

  • context_data (PyTree[Var] | None) – Conditioning data (a pytree) sharing context_encoders’ structure; each leaf is fed to its encoder. None is fully unconditional.

  • context_mask (PyTree[Mask | None] | None) – A pytree of masks with the same structure as context_data. A None leaf (or whole None) is fully conditional; a 0 leaf zeros that source’s contribution.

Returns:

The summed context vector consumed by the backbone and decoders.

Raises:
  • ValueError – If context_data’s structure does not match context_encoders (e.g. a missing/extra source). Whole-tree None is fine.

  • ValueError – If context_data is passed but context_encoders is None: likely user error (nothing would consume the data).

  • ValueError – If context_mask is passed but context_data is None: mask without data is nonsensical.

Return type:

Var