Context encoders¶
- class stix.nn.SumContextEncoder(*args, **kwargs)¶
Bases:
ModuleSum per-source context vectors into a single context.
Pipeline:
time_encoder(t)plus, for each encoder incontext_encoders, itsjax.tree.mapover the matching leaf ofcontext_data, gated by the matching leaf ofcontext_mask; all terms are then summed.context_encodersis an arbitrary pytree of modules andcontext_data/context_maskmust 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 byget_classifier_free_guidance_generator()for its unconditional branch.Per-leaf zero: a
context_maskleaf of0with any (real-valued)context_dataleaf — 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_maskleaf ofNone(or wholeNone) 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.Noneis fully unconditional.context_mask (PyTree[Mask | None] | None) – A pytree of masks with the same structure as
context_data. ANoneleaf (or wholeNone) is fully conditional; a0leaf 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 matchcontext_encoders(e.g. a missing/extra source). Whole-treeNoneis fine.ValueError – If
context_datais passed butcontext_encodersisNone: likely user error (nothing would consume the data).ValueError – If
context_maskis passed butcontext_dataisNone: mask without data is nonsensical.
- Return type: