nur-dev's picture
Add files using upload-large-folder tool
7c5e40e verified
Raw
History Blame Contribute Delete
1.69 kB
"""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)