"""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)