nur-dev's picture
Add files using upload-large-folder tool
7c5e40e verified
Raw
History Blame Contribute Delete
5.39 kB
"""Predicted bounded spans over completed causal segments."""
from __future__ import annotations
from dataclasses import dataclass
import torch
from torch import nn
from strata.modeling.ph_pat.config import PHPATConfig
from strata.modeling.ph_pat.segment_commit import SegmentLayout
@dataclass(frozen=True, slots=True)
class SpanCompilerOutput:
nodes: torch.Tensor
scores: torch.Tensor
starts: torch.Tensor
ends: torch.Tensor
segment_ids: torch.Tensor
valid_mask: torch.Tensor
boundary_logits: torch.Tensor
class SpanCompiler(nn.Module):
def __init__(self, config: PHPATConfig) -> None:
super().__init__()
self.config = config
self.start = nn.Linear(config.d_model, config.d_model, bias=False)
self.end = nn.Linear(config.d_model, config.d_model, bias=False)
self.pool = nn.Linear(config.d_model, config.d_model, bias=False)
self.score = nn.Sequential(
nn.Linear(3 * config.d_model, config.d_model),
nn.SiLU(),
nn.Linear(config.d_model, 1),
)
self.boundary = nn.Linear(config.d_model, 2)
relative_starts, relative_ends = _candidate_bounds(0, config.segment_size, config.max_span_width, torch.device("cpu"))
self.register_buffer("relative_starts", relative_starts, persistent=False)
self.register_buffer("relative_ends", relative_ends, persistent=False)
def forward(self, hidden: torch.Tensor, layout: SegmentLayout) -> SpanCompilerOutput:
batch, _seq_len, dim = hidden.shape
per_segment = self.config.spans_per_segment
prefix = torch.cat((hidden.new_zeros(batch, 1, dim), hidden.cumsum(dim=1)), dim=1)
boundary_logits = self.boundary(hidden)
segment_offsets = layout.starts.view(1, layout.segment_count, 1)
candidate_starts = segment_offsets + self.relative_starts.view(1, 1, -1)
candidate_ends = segment_offsets + self.relative_ends.view(1, 1, -1)
in_sequence = candidate_ends < hidden.shape[1]
safe_starts = candidate_starts.clamp_max(hidden.shape[1] - 1)
safe_ends = candidate_ends.clamp_max(hidden.shape[1] - 1)
row = torch.arange(batch, device=hidden.device).view(batch, 1, 1)
start_score = boundary_logits[row, safe_starts.expand(batch, -1, -1), 0]
end_score = boundary_logits[row, safe_ends.expand(batch, -1, -1), 1]
candidate_valid = layout.segment_valid.unsqueeze(-1) & in_sequence
candidate_scores = (start_score + end_score).masked_fill(~candidate_valid, torch.finfo(hidden.dtype).min)
take = min(per_segment, candidate_scores.shape[-1])
top_scores, top_idx = torch.topk(candidate_scores, k=take, dim=-1)
expanded_starts = candidate_starts.expand(batch, -1, -1)
expanded_ends = candidate_ends.expand(batch, -1, -1)
top_starts = expanded_starts.gather(2, top_idx).clamp_max(hidden.shape[1] - 1)
top_ends = expanded_ends.gather(2, top_idx).clamp_max(hidden.shape[1] - 1)
valid = candidate_valid.expand(batch, -1, -1).gather(2, top_idx)
batch_index = torch.arange(batch, device=hidden.device).view(batch, 1, 1)
span_sum = prefix[batch_index, top_ends + 1] - prefix[batch_index, top_starts]
pooled = span_sum / (top_ends - top_starts + 1).to(hidden.dtype).unsqueeze(-1)
start_hidden = hidden[batch_index, top_starts]
end_hidden = hidden[batch_index, top_ends]
top_nodes = (self.start(start_hidden) + self.end(end_hidden) + self.pool(pooled)) / 3.0
selected_features = torch.cat((start_hidden, end_hidden, pooled), dim=-1)
top_scores = top_scores + self.score(selected_features).squeeze(-1)
if take < per_segment:
pad = per_segment - take
top_nodes = torch.cat((top_nodes, hidden.new_zeros(batch, layout.segment_count, pad, dim)), dim=2)
top_scores = torch.cat((top_scores, hidden.new_full((batch, layout.segment_count, pad), torch.finfo(hidden.dtype).min)), dim=2)
top_starts = torch.cat((top_starts, torch.zeros(batch, layout.segment_count, pad, device=hidden.device, dtype=torch.long)), dim=2)
top_ends = torch.cat((top_ends, torch.zeros(batch, layout.segment_count, pad, device=hidden.device, dtype=torch.long)), dim=2)
valid = torch.cat((valid, torch.zeros(batch, layout.segment_count, pad, device=hidden.device, dtype=torch.bool)), dim=2)
segment_ids = torch.arange(layout.segment_count, device=hidden.device).view(1, -1, 1).expand(batch, -1, per_segment)
return SpanCompilerOutput(
nodes=top_nodes.flatten(1, 2),
scores=top_scores.flatten(1, 2),
starts=top_starts.flatten(1, 2),
ends=top_ends.flatten(1, 2),
segment_ids=segment_ids.flatten(1, 2),
valid_mask=valid.flatten(1, 2),
boundary_logits=boundary_logits,
)
def _candidate_bounds(start: int, end: int, max_width: int, device: torch.device) -> tuple[torch.Tensor, torch.Tensor]:
starts: list[int] = []
ends: list[int] = []
for left in range(start, end):
for right in range(left, min(end, left + max_width)):
starts.append(left)
ends.append(right)
return torch.tensor(starts, device=device), torch.tensor(ends, device=device)
__all__ = ["SpanCompiler", "SpanCompilerOutput"]