CoolFace
Modelpublic

violetxi/hparam-92m-block-mtp-truncated-peak_lr-3e-4-wu-0p01-s42

sourceHugging Faceupdated 17d agoView on Hugging Face
0likes385downloads
native_layers.py337 linesDownload Raw Back to root
1"""Portable HRM layers adapted from Sapient Intelligence's official code.2 3Upstream: https://github.com/sapientinc/HRM4Pinned revision: ac15626f8db096a63c775b84c9dc868776a6feda5 6The original implementation is Apache-2.0 licensed. This derivative replaces7the mandatory FlashAttention import with automatic FlashAttention/PyTorch SDPA8selection and adds padding-mask support. See ``third_party/sapient_hrm``.9"""10 11from __future__ import annotations12 13import importlib14import math15from collections.abc import Callable16from functools import lru_cache17 18import torch19import torch.nn.functional as F20from torch import nn21 22CosSin = tuple[torch.Tensor, torch.Tensor]23 24 25def trunc_normal_init_(26    tensor: torch.Tensor,27    std: float = 1.0,28    lower: float = -2.0,29    upper: float = 2.0,30) -> torch.Tensor:31    """JAX-style truncated normal used by the official HRM release."""32 33    with torch.no_grad():34        if std == 0:35            tensor.zero_()36        else:37            sqrt2 = math.sqrt(2)38            a = math.erf(lower / sqrt2)39            b = math.erf(upper / sqrt2)40            z = (b - a) / 241            c = (2 * math.pi) ** -0.542            pdf_u = c * math.exp(-0.5 * lower**2)43            pdf_l = c * math.exp(-0.5 * upper**2)44            comp_std = std / math.sqrt(45                146                - (upper * pdf_u - lower * pdf_l) / z47                - ((pdf_u - pdf_l) / z) ** 248            )49            tensor.uniform_(a, b)50            tensor.erfinv_()51            tensor.mul_(sqrt2 * comp_std)52            tensor.clip_(lower * comp_std, upper * comp_std)53    return tensor54 55 56class CastedLinear(nn.Module):57    def __init__(self, in_features: int, out_features: int, bias: bool) -> None:58        super().__init__()59        self.weight = nn.Parameter(60            trunc_normal_init_(61                torch.empty(out_features, in_features),62                std=1.0 / math.sqrt(in_features),63            )64        )65        self.bias = nn.Parameter(torch.zeros(out_features)) if bias else None66 67    def forward(self, inputs: torch.Tensor) -> torch.Tensor:68        bias = self.bias.to(inputs.dtype) if self.bias is not None else None69        return F.linear(inputs, self.weight.to(inputs.dtype), bias)70 71 72class CastedEmbedding(nn.Module):73    def __init__(74        self,75        num_embeddings: int,76        embedding_dim: int,77        init_std: float,78        cast_to: torch.dtype,79    ) -> None:80        super().__init__()81        self.cast_to = cast_to82        self.embedding_weight = nn.Parameter(83            trunc_normal_init_(84                torch.empty(num_embeddings, embedding_dim), std=init_std85            )86        )87 88    def forward(self, inputs: torch.Tensor) -> torch.Tensor:89        return F.embedding(inputs.to(torch.long), self.embedding_weight.to(self.cast_to))90 91 92def rotate_half(inputs: torch.Tensor) -> torch.Tensor:93    first, second = inputs.chunk(2, dim=-1)94    return torch.cat((-second, first), dim=-1)95 96 97def apply_rotary_pos_emb(98    query: torch.Tensor,99    key: torch.Tensor,100    cos: torch.Tensor,101    sin: torch.Tensor,102) -> tuple[torch.Tensor, torch.Tensor]:103    original_dtype = query.dtype104    query_f32 = query.to(cos.dtype)105    key_f32 = key.to(cos.dtype)106    cos = cos.unsqueeze(-2)107    sin = sin.unsqueeze(-2)108    query = query_f32 * cos + rotate_half(query_f32) * sin109    key = key_f32 * cos + rotate_half(key_f32) * sin110    return query.to(original_dtype), key.to(original_dtype)111 112 113class RotaryEmbedding(nn.Module):114    def __init__(self, dim: int, max_position_embeddings: int, base: float) -> None:115        super().__init__()116        inv_freq = 1.0 / (117            base ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim)118        )119        positions = torch.arange(max_position_embeddings, dtype=torch.float32)120        frequencies = torch.outer(positions, inv_freq)121        embedding = torch.cat((frequencies, frequencies), dim=-1)122        self.register_buffer("cos_cached", embedding.cos(), persistent=False)123        self.register_buffer("sin_cached", embedding.sin(), persistent=False)124 125    def forward(self, seq_len: int) -> CosSin:126        return self.cos_cached[:seq_len], self.sin_cached[:seq_len]127 128    def for_positions(self, position_ids: torch.Tensor) -> CosSin:129        """Return rotary values for an explicit one-dimensional position map."""130 131        if position_ids.ndim != 1:132            raise ValueError("position_ids must have shape [sequence_length]")133        position_ids = position_ids.to(134            device=self.cos_cached.device, dtype=torch.long135        )136        if bool(137            ((position_ids < 0) | (position_ids >= self.cos_cached.shape[0])).any()138        ):139            raise ValueError("position_ids contain an out-of-range position")140        return (141            self.cos_cached.index_select(0, position_ids),142            self.sin_cached.index_select(0, position_ids),143        )144 145 146@lru_cache(maxsize=1)147def _flash_attention_function() -> Callable[..., torch.Tensor] | None:148    for module_name in ("flash_attn_interface", "flash_attn"):149        try:150            module = importlib.import_module(module_name)151        except (ImportError, OSError):152            continue153        function = getattr(module, "flash_attn_func", None)154        if function is not None:155            return function156    return None157 158 159def flash_attention_available() -> bool:160    return _flash_attention_function() is not None161 162 163class Attention(nn.Module):164    def __init__(165        self,166        hidden_size: int,167        num_heads: int,168        backend: str = "auto",169        causal: bool = False,170    ) -> None:171        super().__init__()172        self.hidden_size = hidden_size173        self.num_heads = num_heads174        self.head_dim = hidden_size // num_heads175        self.backend = backend176        self.causal = bool(causal)177        self.last_backend = "uninitialized"178        self.qkv_proj = CastedLinear(hidden_size, 3 * hidden_size, bias=False)179        self.o_proj = CastedLinear(hidden_size, hidden_size, bias=False)180 181    def _flash_supported(182        self,183        hidden_states: torch.Tensor,184        attention_mask: torch.Tensor | None,185        attention_pattern: torch.Tensor | None,186    ) -> bool:187        if hidden_states.device.type != "cuda":188            return False189        if hidden_states.dtype not in (torch.float16, torch.bfloat16):190            return False191        if self.head_dim > 256:192            return False193        if attention_mask is not None and not bool(attention_mask.all().item()):194            return False195        if attention_pattern is not None:196            return False197        return _flash_attention_function() is not None198 199    def _sdpa(200        self,201        query: torch.Tensor,202        key: torch.Tensor,203        value: torch.Tensor,204        attention_mask: torch.Tensor | None,205        attention_pattern: torch.Tensor | None,206    ) -> torch.Tensor:207        query = query.transpose(1, 2)208        key = key.transpose(1, 2)209        value = value.transpose(1, 2)210        sdpa_mask = None211        is_causal = self.causal212        if attention_pattern is not None:213            query_length = query.shape[-2]214            key_length = key.shape[-2]215            if attention_pattern.ndim == 2:216                if attention_pattern.shape != (query_length, key_length):217                    raise ValueError(218                        "attention_pattern must have shape [query_length, key_length]"219                    )220                sdpa_mask = attention_pattern[None, None]221            elif attention_pattern.ndim == 3:222                if attention_pattern.shape != (223                    query.shape[0],224                    query_length,225                    key_length,226                ):227                    raise ValueError(228                        "batched attention_pattern must have shape "229                        "[batch_size, query_length, key_length]"230                    )231                sdpa_mask = attention_pattern[:, None]232            else:233                raise ValueError("attention_pattern must have rank two or three")234            sdpa_mask = sdpa_mask.to(device=query.device, dtype=torch.bool)235            # A structural pattern is authoritative. Parallel lanes are236            # physically appended after the source but may attend only to237            # their own gold prefix and their own recurrent lane state.238            is_causal = False239 240        if attention_mask is not None and not bool(attention_mask.all().item()):241            key_mask = attention_mask[:, None, None, :].to(242                device=query.device, dtype=torch.bool243            )244            sdpa_mask = key_mask if sdpa_mask is None else sdpa_mask & key_mask245            if self.causal and attention_pattern is None:246                sequence_length = query.shape[-2]247                causal_mask = torch.ones(248                    (sequence_length, sequence_length),249                    dtype=torch.bool,250                    device=query.device,251                ).tril_()252                sdpa_mask = sdpa_mask & causal_mask[None, None, :, :]253                is_causal = False254        output = F.scaled_dot_product_attention(255            query,256            key,257            value,258            attn_mask=sdpa_mask,259            dropout_p=0.0,260            is_causal=is_causal,261        )262        return output.transpose(1, 2)263 264    def forward(265        self,266        cos_sin: CosSin,267        hidden_states: torch.Tensor,268        attention_mask: torch.Tensor | None = None,269        attention_pattern: torch.Tensor | None = None,270    ) -> torch.Tensor:271        batch_size, seq_len, _ = hidden_states.shape272        qkv = self.qkv_proj(hidden_states).view(273            batch_size, seq_len, 3, self.num_heads, self.head_dim274        )275        query, key, value = qkv.unbind(dim=2)276        query, key = apply_rotary_pos_emb(query, key, *cos_sin)277 278        can_flash = self._flash_supported(279            hidden_states, attention_mask, attention_pattern280        )281        if self.backend == "flash" and not can_flash:282            raise RuntimeError(283                "FlashAttention was requested but is unavailable or unsupported for "284                "this device, dtype, head dimension, or padding mask"285            )286 287        if self.backend in {"auto", "flash"} and can_flash:288            flash = _flash_attention_function()289            assert flash is not None290            try:291                output = flash(292                    query,293                    key,294                    value,295                    dropout_p=0.0,296                    causal=self.causal,297                )298                if isinstance(output, tuple):299                    output = output[0]300                self.last_backend = "flash"301            except RuntimeError:302                if self.backend == "flash":303                    raise304                output = self._sdpa(305                    query, key, value, attention_mask, attention_pattern306                )307                self.last_backend = "sdpa"308        else:309            output = self._sdpa(310                query, key, value, attention_mask, attention_pattern311            )312            self.last_backend = "sdpa"313 314        output = output.reshape(batch_size, seq_len, self.hidden_size)315        return self.o_proj(output)316 317 318class SwiGLU(nn.Module):319    def __init__(self, hidden_size: int, intermediate_size: int) -> None:320        super().__init__()321        self.gate_up_proj = CastedLinear(322            hidden_size, 2 * intermediate_size, bias=False323        )324        self.down_proj = CastedLinear(intermediate_size, hidden_size, bias=False)325 326    def forward(self, inputs: torch.Tensor) -> torch.Tensor:327        gate, up = self.gate_up_proj(inputs).chunk(2, dim=-1)328        return self.down_proj(F.silu(gate) * up)329 330 331def rms_norm(hidden_states: torch.Tensor, variance_epsilon: float) -> torch.Tensor:332    input_dtype = hidden_states.dtype333    hidden_states = hidden_states.to(torch.float32)334    variance = hidden_states.square().mean(-1, keepdim=True)335    hidden_states = hidden_states * torch.rsqrt(variance + variance_epsilon)336    return hidden_states.to(input_dtype)337