| """Attention building blocks used by RiboSphere.""" |
|
|
| from __future__ import annotations |
|
|
| from typing import Any |
|
|
| import torch |
| from torch import Tensor, nn |
| import torch.nn.functional as F |
| from torch.nn.attention.flex_attention import create_block_mask, flex_attention |
|
|
| from .layers import FeedForward |
| from .rotary import RotaryEmbedding |
|
|
| AttentionArguments = dict[str, Any] |
|
|
|
|
| def root_mean_square_norm(tensor: Tensor) -> Tensor: |
| """Apply parameter-free RMS normalization over the final dimension.""" |
|
|
| return F.rms_norm(tensor, (tensor.shape[-1],)) |
|
|
|
|
| class TransformerStack(nn.Module): |
| """Stack of local self-attention blocks.""" |
|
|
| def __init__( |
| self, |
| *, |
| num_channels: int, |
| num_heads: int, |
| mlp_factor: int, |
| window_size: int, |
| num_layers: int, |
| attention_backend: str = "flex", |
| dropout: float = 0.1, |
| pairwise_channels: int = 0, |
| is_causal: bool = False, |
| ) -> None: |
| super().__init__() |
| if num_channels <= 0 or num_heads <= 0 or num_layers <= 0: |
| raise ValueError( |
| "num_channels, num_heads, and num_layers must be positive." |
| ) |
| if num_channels % num_heads != 0: |
| raise ValueError("num_channels must be divisible by num_heads.") |
| if window_size <= 0: |
| raise ValueError("window_size must be positive.") |
| if pairwise_channels < 0: |
| raise ValueError("pairwise_channels cannot be negative.") |
|
|
| attention_backend = attention_backend.lower() |
| if attention_backend not in {"sdpa", "flex"}: |
| raise ValueError("attention_backend must be 'sdpa' or 'flex'.") |
|
|
| use_pair_bias = pairwise_channels > 0 |
| self.blocks = nn.ModuleList( |
| [ |
| TransformerBlock( |
| num_channels=num_channels, |
| num_heads=num_heads, |
| mlp_factor=mlp_factor, |
| attention_backend=attention_backend, |
| dropout=dropout, |
| use_pairwise_bias=use_pair_bias, |
| pairwise_channels=pairwise_channels, |
| ) |
| for _ in range(num_layers) |
| ] |
| ) |
| self.window_size = window_size |
| self.is_causal = is_causal |
| self.attention_backend = attention_backend |
|
|
| def _window_mask( |
| self, |
| batch_index: Tensor, |
| head_index: Tensor, |
| query_index: Tensor, |
| key_value_index: Tensor, |
| ) -> Tensor: |
| del batch_index, head_index |
| within_window = (query_index - key_value_index).abs() <= self.window_size |
| if self.is_causal: |
| within_window = within_window & (query_index >= key_value_index) |
| return within_window |
|
|
| def forward( |
| self, |
| hidden_states: Tensor, |
| pairwise_features: Tensor | None = None, |
| ) -> Tensor: |
| """Transform ``[B, L, D]`` token features.""" |
|
|
| if hidden_states.ndim != 3: |
| raise ValueError("hidden_states must have shape [B, L, D].") |
|
|
| sequence_length = hidden_states.shape[1] |
| if pairwise_features is not None and pairwise_features.shape[:3] != ( |
| hidden_states.shape[0], |
| sequence_length, |
| sequence_length, |
| ): |
| raise ValueError( |
| "pairwise_features must have shape [B, L, L, P]." |
| ) |
|
|
| if self.attention_backend == "flex": |
| attention_arguments: AttentionArguments = { |
| "block_mask": create_block_mask( |
| self._window_mask, |
| B=None, |
| H=None, |
| Q_LEN=sequence_length, |
| KV_LEN=sequence_length, |
| device=hidden_states.device, |
| ), |
| "score_mod": None, |
| } |
| else: |
| positions = torch.arange( |
| sequence_length, |
| device=hidden_states.device, |
| ) |
| attention_mask = ( |
| positions[:, None] - positions[None, :] |
| ).abs() <= self.window_size |
| if self.is_causal: |
| attention_mask = attention_mask & ( |
| positions[:, None] >= positions[None, :] |
| ) |
| attention_arguments = {"attn_mask": attention_mask.unsqueeze(0)} |
|
|
| for block in self.blocks: |
| hidden_states = block( |
| hidden_states, |
| pairwise_features=pairwise_features, |
| **attention_arguments, |
| ) |
| return hidden_states |
|
|
|
|
| class TransformerBlock(nn.Module): |
| """Pre-normalized self-attention and feed-forward block.""" |
|
|
| def __init__( |
| self, |
| *, |
| num_channels: int, |
| num_heads: int, |
| mlp_factor: int, |
| attention_backend: str = "flex", |
| dropout: float = 0.1, |
| use_pairwise_bias: bool = False, |
| pairwise_channels: int = 0, |
| ) -> None: |
| super().__init__() |
| self.attention_backend = attention_backend |
| self.attention = SelfAttention( |
| model_dimension=num_channels, |
| num_heads=num_heads, |
| dropout=dropout, |
| attention_backend=attention_backend, |
| ) |
| self.feed_forward = FeedForward( |
| num_channels, |
| num_channels * mlp_factor, |
| num_channels, |
| activation=nn.GELU, |
| dropout=dropout, |
| ) |
|
|
| self.use_pairwise_bias = use_pairwise_bias |
| if use_pairwise_bias: |
| if pairwise_channels <= 0: |
| raise ValueError( |
| "pairwise_channels must be positive when pair bias is enabled." |
| ) |
| self.pair_bias_projection = nn.Linear( |
| pairwise_channels, 1, bias=True |
| ) |
| self.pair_bias_norm = nn.LayerNorm(pairwise_channels) |
| else: |
| self.pair_bias_projection = None |
| self.pair_bias_norm = None |
|
|
| def _add_pair_bias( |
| self, |
| pairwise_features: Tensor, |
| attention_arguments: AttentionArguments, |
| ) -> AttentionArguments: |
| if self.pair_bias_projection is None or self.pair_bias_norm is None: |
| return attention_arguments |
|
|
| pair_bias = self.pair_bias_projection( |
| self.pair_bias_norm(pairwise_features) |
| ).squeeze(-1) |
| attention_arguments = dict(attention_arguments) |
|
|
| if self.attention_backend == "flex": |
|
|
| def pair_biased_score( |
| score: Tensor, |
| batch_index: Tensor, |
| head_index: Tensor, |
| query_index: Tensor, |
| key_value_index: Tensor, |
| ) -> Tensor: |
| del head_index |
| return score + pair_bias[ |
| batch_index, |
| query_index, |
| key_value_index, |
| ] |
|
|
| attention_arguments["score_mod"] = pair_biased_score |
| else: |
| attention_mask = attention_arguments["attn_mask"] |
| additive_pair_bias = torch.where( |
| attention_mask, |
| pair_bias, |
| torch.full_like(pair_bias, -torch.inf), |
| ) |
| attention_arguments["attn_mask"] = additive_pair_bias.unsqueeze(1) |
| return attention_arguments |
|
|
| def forward( |
| self, |
| hidden_states: Tensor, |
| *, |
| pairwise_features: Tensor | None = None, |
| **attention_arguments: Any, |
| ) -> Tensor: |
| if self.use_pairwise_bias: |
| if pairwise_features is None: |
| raise ValueError( |
| "pairwise_features are required when pair bias is enabled." |
| ) |
| attention_arguments = self._add_pair_bias( |
| pairwise_features, |
| attention_arguments, |
| ) |
|
|
| hidden_states = hidden_states + self.attention( |
| root_mean_square_norm(hidden_states), |
| **attention_arguments, |
| ) |
| hidden_states = hidden_states + self.feed_forward( |
| root_mean_square_norm(hidden_states) |
| ) |
| return hidden_states |
|
|
|
|
| class SelfAttention(nn.Module): |
| """Multi-head self-attention with rotary position embeddings.""" |
|
|
| def __init__( |
| self, |
| model_dimension: int, |
| num_heads: int, |
| *, |
| normalize_queries_and_keys: bool = False, |
| attention_backend: str = "flex", |
| dropout: float = 0.1, |
| ) -> None: |
| super().__init__() |
| if model_dimension <= 0 or num_heads <= 0: |
| raise ValueError("model_dimension and num_heads must be positive.") |
| if model_dimension % num_heads != 0: |
| raise ValueError("model_dimension must be divisible by num_heads.") |
| if not 0.0 <= dropout < 1.0: |
| raise ValueError("dropout must be in [0, 1).") |
|
|
| attention_backend = attention_backend.lower() |
| if attention_backend not in {"flex", "sdpa"}: |
| raise ValueError("backend must be 'flex' or 'sdpa'.") |
|
|
| self.model_dimension = model_dimension |
| self.num_heads = num_heads |
| self.head_dimension = self.model_dimension // self.num_heads |
| self.attention_dropout = nn.Dropout(dropout) |
| self.dropout = dropout |
| self.normalize_queries_and_keys = normalize_queries_and_keys |
|
|
| self.rotary_embedding = RotaryEmbedding(self.head_dimension) |
| self.qkv_projection = nn.Linear( |
| model_dimension, 3 * model_dimension, bias=True |
| ) |
| self.output_projection = nn.Linear(model_dimension, model_dimension) |
| self.residual_dropout = nn.Dropout(dropout) |
| self.attention_backend = attention_backend |
|
|
| def forward( |
| self, |
| hidden_states: Tensor, |
| **attention_arguments: Any, |
| ) -> Tensor: |
| """Apply self-attention to ``[B, L, D]`` hidden states.""" |
|
|
| if hidden_states.ndim != 3: |
| raise ValueError("hidden_states must have shape [B, L, D].") |
| batch_size, sequence_length, hidden_dimension = hidden_states.shape |
| if hidden_dimension != self.model_dimension: |
| raise ValueError( |
| f"Expected hidden dimension {self.model_dimension}, " |
| f"received {hidden_dimension}." |
| ) |
|
|
| query, key, value = self.qkv_projection(hidden_states).split( |
| self.model_dimension, dim=-1 |
| ) |
|
|
| def split_heads(tensor: Tensor) -> Tensor: |
| return tensor.reshape( |
| batch_size, |
| sequence_length, |
| self.num_heads, |
| self.head_dimension, |
| ).transpose(1, 2) |
|
|
| query, key, value = map(split_heads, (query, key, value)) |
| if self.normalize_queries_and_keys: |
| query = root_mean_square_norm(query) |
| key = root_mean_square_norm(key) |
| query, key = self.rotary_embedding(query, key) |
|
|
| if self.attention_backend == "flex": |
| attention_output = flex_attention( |
| query, |
| key, |
| value, |
| block_mask=attention_arguments.get("block_mask"), |
| score_mod=attention_arguments.get("score_mod"), |
| ) |
| else: |
| attention_output = F.scaled_dot_product_attention( |
| query, |
| key, |
| value, |
| **attention_arguments, |
| ) |
|
|
| attention_output = self.attention_dropout(attention_output) |
| attention_output = attention_output.transpose(1, 2).contiguous().view( |
| batch_size, |
| sequence_length, |
| self.model_dimension, |
| ) |
| attention_output = self.residual_dropout( |
| self.output_projection(attention_output) |
| ) |
| return attention_output |
|
|