| """Typed outputs for STRATA models.""" |
|
|
| from __future__ import annotations |
|
|
| from dataclasses import dataclass, field |
|
|
| import torch |
|
|
|
|
| @dataclass(slots=True) |
| class PredicateBlockOutput: |
| """Auxiliary predictions emitted by a predicate block. |
| |
| ``node_type_logits``, ``chart_type_logits`` and ``edge_logits`` are prediction |
| heads that do not feed the forward pass; by default they are computed only on |
| the final predicate block, so they are ``None`` on earlier blocks. |
| ``predicate_gate`` is always present (it modulates predicate-memory attention). |
| """ |
|
|
| node_type_logits: torch.Tensor | None |
| chart_type_logits: torch.Tensor | None |
| predicate_gate: torch.Tensor |
| edge_logits: torch.Tensor | None = None |
|
|
|
|
| @dataclass(slots=True) |
| class GraphObjectBlockOutput: |
| """Auxiliary predictions emitted from explicit graph-object edge memory.""" |
|
|
| relation_logits: torch.Tensor |
| src_logits: torch.Tensor |
| dst_logits: torch.Tensor |
|
|
|
|
| @dataclass(slots=True) |
| class StrataModelOutput: |
| """Decoder backbone output.""" |
|
|
| last_hidden_state: torch.Tensor |
| predicate_outputs: tuple[PredicateBlockOutput, ...] |
| attention_replacement_gates: torch.Tensor |
| graph_object_outputs: tuple[GraphObjectBlockOutput, ...] = field(default_factory=tuple) |
|
|
|
|
| @dataclass(slots=True) |
| class StrataCausalLMOutput: |
| """Causal LM output with optional graph/chart auxiliary predictions.""" |
|
|
| logits: torch.Tensor |
| loss: torch.Tensor | None |
| hidden_states: torch.Tensor |
| predicate_outputs: tuple[PredicateBlockOutput, ...] |
| attention_replacement_gates: torch.Tensor |
| graph_object_outputs: tuple[GraphObjectBlockOutput, ...] = field(default_factory=tuple) |
|
|