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