violetxi/hparam-92m-block-mtp-truncated-peak_lr-3e-4-wu-0p01-s42
0385
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 