diff --git a/src/spine/backbones/base.py b/src/spine/backbones/base.py index 5a78b03..add0ad5 100644 --- a/src/spine/backbones/base.py +++ b/src/spine/backbones/base.py @@ -13,11 +13,16 @@ @dataclass class EncodedEvent: - """What every backbone returns.""" + """What every backbone returns. + + ``cls`` is optional: a backbone with no event-level readout leaves it + ``None``, and consumers that need an event embedding (the query encoder) + must guard for it. Per-token pretexts read ``tokens`` and ignore ``cls``. + """ tokens: Tensor # [B, L, D] per-pulse token embeddings token_mask: Tensor # [B, L] bool, True = real pulse (not padding) - cls: Tensor # [B, D] pooled event embedding + cls: Tensor | None = None # [B, D] pooled event embedding, or None class Backbone(nn.Module): diff --git a/src/spine/pretrain/curtain/head.py b/src/spine/pretrain/curtain/head.py index ac8f2ad..bf31a7c 100644 --- a/src/spine/pretrain/curtain/head.py +++ b/src/spine/pretrain/curtain/head.py @@ -84,7 +84,16 @@ def forward(self, query_pos: Tensor, enc: EncodedEvent) -> Tensor: Returns: [B, Q, D] per-query embeddings. + + Raises: + ValueError: If the backbone returned no event-level embedding + (``enc.cls is None``); this encoder requires one. """ + if enc.cls is None: + raise ValueError( + "QueryCrossAttnEncoder requires an event-level `cls` embedding, " + "but the backbone returned cls=None" + ) kv = torch.cat([enc.cls.unsqueeze(1), enc.tokens], dim=1) # [B,1+L,D] ones = torch.ones( enc.token_mask.shape[0], 1, dtype=torch.bool, device=enc.token_mask.device