"""Pre-LN Transformer encoder that routes self-attention through ``flex_attention``. The block-diffusion masks are structured (block-causal over thought slots plus optional key padding), so expressing them as a ``flex_attention`` block mask lets the fused kernel skip fully-masked blocks instead of materializing a dense ``[batch * heads, L, L]`` score tensor. This keeps attention memory and time near-linear in the trace length, which the default ``nn.TransformerEncoder`` math path does not. The layer geometry mirrors ``nn.TransformerEncoderLayer(norm_first=True, activation='gelu')`` so behaviour matches the default path up to floating-point error; the attention core is verified against a dense masked-softmax reference before use. """ from __future__ import annotations import math import os import torch import torch.nn.functional as F from torch import Tensor, nn from torch.nn.attention.flex_attention import BlockMask, create_block_mask, flex_attention # Compiled dynamic=False is fastest for fixed-shape training. Variable-length generation # recompiles per new sequence length, so serving sets MDLM_FLEX_EAGER=1 to run eager and # trade steady-state speed for the absence of per-length recompilation stalls. _FLEX_EAGER = os.environ.get('MDLM_FLEX_EAGER') == '1' _flex_compiled = flex_attention if _FLEX_EAGER else torch.compile(flex_attention, dynamic=False) def build_block_mask( attn_mask: Tensor | None, padding_mask: Tensor | None, batch_size: int, seq_len: int, device: torch.device, ) -> BlockMask | None: """Build a broadcast-over-heads ``BlockMask`` from project boolean masks. ``attn_mask`` follows the src_mask convention (``True`` blocks a key), shaped ``[L, L]`` or ``[batch, L, L]``. ``padding_mask`` follows the src_key_padding_mask convention (``True`` marks padding). Returns ``None`` when neither constrains attention. """ if attn_mask is None and padding_mask is None: return None shared = attn_mask is not None and attn_mask.dim() == 2 def mask_mod(b: Tensor, h: Tensor, q_idx: Tensor, kv_idx: Tensor) -> Tensor: keep = torch.ones_like(q_idx, dtype=torch.bool) if attn_mask is not None: blocked = attn_mask[q_idx, kv_idx] if shared else attn_mask[b, q_idx, kv_idx] keep = keep & ~blocked if padding_mask is not None: keep = keep & ~padding_mask[b, kv_idx] return keep return create_block_mask( mask_mod, batch_size, None, seq_len, seq_len, device=device, _compile=not _FLEX_EAGER ) def flex_self_attention( query: Tensor, key: Tensor, value: Tensor, block_mask: BlockMask | None ) -> Tensor: """Multi-head self-attention over ``[batch, heads, L, head_dim]`` tensors.""" if block_mask is None: return flex_attention(query, key, value) return _flex_compiled(query, key, value, block_mask=block_mask) class FlexEncoderLayer(nn.Module): """Pre-LN Transformer block with a ``flex_attention`` self-attention core.""" def __init__(self, d_model: int, n_heads: int, d_ff: int, dropout: float) -> None: super().__init__() if d_model % n_heads != 0: raise ValueError("d_model must be divisible by n_heads") self.n_heads = n_heads self.head_dim = d_model // n_heads self.q_proj = nn.Linear(d_model, d_model) self.k_proj = nn.Linear(d_model, d_model) self.v_proj = nn.Linear(d_model, d_model) self.out_proj = nn.Linear(d_model, d_model) self.linear1 = nn.Linear(d_model, d_ff) self.linear2 = nn.Linear(d_ff, d_model) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.dropout = nn.Dropout(dropout) def _split_heads(self, projected: Tensor) -> Tensor: batch_size, seq_len, _ = projected.shape return projected.view(batch_size, seq_len, self.n_heads, self.head_dim).transpose(1, 2) def forward(self, hidden: Tensor, block_mask: BlockMask | None) -> Tensor: normed = self.norm1(hidden) query = self._split_heads(self.q_proj(normed)) key = self._split_heads(self.k_proj(normed)) value = self._split_heads(self.v_proj(normed)) attended = flex_self_attention(query, key, value, block_mask) batch_size, _, seq_len, _ = attended.shape attended = attended.transpose(1, 2).reshape(batch_size, seq_len, -1) hidden = hidden + self.dropout(self.out_proj(attended)) normed = self.norm2(hidden) feed_forward = self.linear2(self.dropout(F.gelu(self.linear1(normed)))) return hidden + self.dropout(feed_forward) class FlexEncoder(nn.Module): """Stack of :class:`FlexEncoderLayer` blocks with a final layer norm.""" def __init__( self, d_model: int, n_heads: int, d_ff: int, dropout: float, n_layers: int, activation_checkpointing: bool, ) -> None: super().__init__() self.layers = nn.ModuleList( FlexEncoderLayer(d_model, n_heads, d_ff, dropout) for _ in range(n_layers) ) self.norm = nn.LayerNorm(d_model) self.activation_checkpointing = activation_checkpointing def init_residual_outputs(self, n_layers: int) -> None: residual_std = 0.02 / math.sqrt(2 * n_layers) for layer in self.layers: nn.init.normal_(layer.out_proj.weight, mean=0.0, std=residual_std) nn.init.normal_(layer.linear2.weight, mean=0.0, std=residual_std) def forward(self, hidden: Tensor, block_mask: BlockMask | None) -> Tensor: use_checkpoint = ( self.activation_checkpointing and self.training and torch.is_grad_enabled() ) for layer in self.layers: if use_checkpoint: hidden = torch.utils.checkpoint.checkpoint( layer, hidden, block_mask, use_reentrant=False ) else: hidden = layer(hidden, block_mask) return self.norm(hidden)