File size: 2,399 Bytes
e69b72a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 | """Predicted dependency and constituent structure over induced spans."""
from __future__ import annotations
from dataclasses import dataclass
import math
import torch
from torch import nn
from strata.modeling.ph_pat.config import PHPATConfig
from strata.modeling.ph_pat.span_compiler import SpanCompilerOutput
@dataclass(frozen=True, slots=True)
class DependencyChartOutput:
nodes: torch.Tensor
head_logits: torch.Tensor
relation_logits: torch.Tensor
constituent_logits: torch.Tensor
segment_ids: torch.Tensor
valid_mask: torch.Tensor
class DependencyChart(nn.Module):
def __init__(self, config: PHPATConfig) -> None:
super().__init__()
self.config = config
self.head_query = nn.Linear(config.d_model, config.d_model, bias=False)
self.head_key = nn.Linear(config.d_model, config.d_model, bias=False)
self.relation = nn.Linear(2 * config.d_model, config.dependency_relations)
self.constituent = nn.Linear(config.d_model, config.constituent_categories)
self.node_update = nn.Sequential(
nn.Linear(config.d_model, config.d_model),
nn.SiLU(),
nn.Linear(config.d_model, config.d_model),
)
def forward(self, spans: SpanCompilerOutput) -> DependencyChartOutput:
nodes = spans.nodes + self.node_update(spans.nodes)
query = self.head_query(nodes)
key = self.head_key(nodes)
logits = torch.matmul(query, key.transpose(-2, -1)) / math.sqrt(self.config.d_model)
same_segment = spans.segment_ids.unsqueeze(-1) == spans.segment_ids.unsqueeze(-2)
valid_pair = spans.valid_mask.unsqueeze(-1) & spans.valid_mask.unsqueeze(-2) & same_segment
logits = logits.masked_fill(~valid_pair, torch.finfo(logits.dtype).min)
head_weights = torch.softmax(logits.float(), dim=-1).to(nodes.dtype)
head_weights = torch.nan_to_num(head_weights)
parent = torch.matmul(head_weights, nodes)
relation_logits = self.relation(torch.cat((nodes, parent), dim=-1))
return DependencyChartOutput(
nodes=nodes,
head_logits=logits,
relation_logits=relation_logits,
constituent_logits=self.constituent(nodes),
segment_ids=spans.segment_ids,
valid_mask=spans.valid_mask,
)
__all__ = ["DependencyChart", "DependencyChartOutput"]
|