Download src/diffusion_lm/flexattn.py from goldenfox/marimo-diffusion: direct link, hf CLI and curl.
- Browser
- Download file 6.06 kB
-
https://huggingface.co/goldenfox/marimo-diffusion/resolve/main/src/diffusion_lm/flexattn.py
- Command line
-
hf download hf://goldenfox/marimo-diffusion/src/diffusion_lm/flexattn.py
-
curl -L -o flexattn.py https://huggingface.co/goldenfox/marimo-diffusion/resolve/main/src/diffusion_lm/flexattn.py
6.06 kB
| """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) | |