import math import torch import torch.nn as nn import torch.nn.functional as F from typing import Optional, Tuple from einops import repeat from einops.layers.torch import Rearrange # ----------------------------------------------------------------------------- # small utils (keep yours if already present) # ----------------------------------------------------------------------------- def exists(val): return val is not None def rand_bool(shape, proba, device=None): if proba == 1: return torch.ones(shape, device=device, dtype=torch.bool) if proba == 0: return torch.zeros(shape, device=device, dtype=torch.bool) return torch.bernoulli(torch.full(shape, proba, device=device)).to(torch.bool) # ----------------------------------------------------------------------------- # time embedding (reuse your existing ones if already in this file) # ----------------------------------------------------------------------------- class LearnedPositionalEmbedding(nn.Module): def __init__(self, dim: int): super().__init__() assert (dim % 2) == 0 half_dim = dim // 2 self.weights = nn.Parameter(torch.randn(half_dim)) def forward(self, x: torch.Tensor) -> torch.Tensor: # x: [B] x = x.view(-1, 1) freqs = x * self.weights.view(1, -1) * 2 * math.pi fouriered = torch.cat((freqs.sin(), freqs.cos()), dim=-1) return torch.cat((x, fouriered), dim=-1) def TimePositionalEmbedding(dim: int, out_features: int) -> nn.Module: return nn.Sequential( LearnedPositionalEmbedding(dim), nn.Linear(dim + 1, out_features), ) class FixedEmbedding(nn.Module): """ Safer than your current version: supports any length by wrapping positions. """ def __init__(self, max_length: int, features: int): super().__init__() self.max_length = max_length self.embedding = nn.Embedding(max_length, features) def forward(self, x: torch.Tensor) -> torch.Tensor: b, t = x.shape[0], x.shape[1] device = x.device pos = torch.arange(t, device=device) % self.max_length emb = self.embedding(pos) # [T, D] return repeat(emb, "t d -> b t d", b=b) # ----------------------------------------------------------------------------- # HRM building blocks (RoPE attention + SwiGLU FFN + AdaRMSNorm) # ----------------------------------------------------------------------------- class AdaRMSNorm(nn.Module): def __init__(self, style_dim: int, channels: int, eps: float = 1e-5): super().__init__() self.fc = nn.Linear(style_dim, channels * 2) self.eps = eps with torch.no_grad(): self.fc.weight.zero_() self.fc.bias.zero_() def forward(self, x: torch.Tensor, style: torch.Tensor) -> torch.Tensor: # x: [B, T, C], style: [B, S] gb = self.fc(style) gamma, beta = gb.chunk(2, dim=-1) gamma = gamma.unsqueeze(1).to(dtype=x.dtype) beta = beta.unsqueeze(1).to(dtype=x.dtype) x_norm = x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps) return (1.0 + gamma) * x_norm + beta class SwiGLUFFN(nn.Module): def __init__(self, dim: int, expansion: float = 2.0, dropout: float = 0.0): super().__init__() hidden = int(dim * expansion) self.fc = nn.Linear(dim, hidden * 2) self.proj = nn.Linear(hidden, dim) self.dropout = nn.Dropout(dropout) def forward(self, x: torch.Tensor) -> torch.Tensor: a, b = self.fc(x).chunk(2, dim=-1) x = F.silu(b) * a x = self.proj(self.dropout(x)) return x class RotaryEmbedding(nn.Module): def __init__(self, dim: int, base: float = 10000.0): super().__init__() if dim % 2 != 0: raise ValueError("RoPE head_dim must be even.") inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2).float() / dim)) self.register_buffer("inv_freq", inv_freq, persistent=False) def get_cos_sin( self, seq_len: int, device: torch.device, dtype: torch.dtype, ) -> Tuple[torch.Tensor, torch.Tensor]: t = torch.arange(seq_len, device=device, dtype=dtype) freqs = torch.outer(t, self.inv_freq.to(dtype)) emb = torch.cat([freqs, freqs], dim=-1) # [T, head_dim] cos = emb.cos().unsqueeze(0).unsqueeze(0) # [1, 1, T, D] sin = emb.sin().unsqueeze(0).unsqueeze(0) return cos, sin def apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: # x: [B, H, T, D] x_even = x[..., ::2] x_odd = x[..., 1::2] x_rot = torch.stack([-x_odd, x_even], dim=-1).flatten(-2) return (x * cos) + (x_rot * sin) class MultiheadSelfAttentionRoPE(nn.Module): def __init__( self, embed_dim: int, num_heads: int, dropout: float = 0.0, use_rope: bool = True, rope_base: float = 10000.0, ): super().__init__() if embed_dim % num_heads != 0: raise ValueError("embed_dim must be divisible by num_heads.") self.embed_dim = embed_dim self.num_heads = num_heads self.head_dim = embed_dim // num_heads self.scale = 1.0 / math.sqrt(self.head_dim) self.qkv = nn.Linear(embed_dim, embed_dim * 3) self.out_proj = nn.Linear(embed_dim, embed_dim) self.attn_dropout = nn.Dropout(dropout) self.use_rope = use_rope self.rope = RotaryEmbedding(self.head_dim, base=rope_base) if use_rope else None if use_rope and (self.head_dim % 2 != 0): raise ValueError( f"RoPE requires even head_dim, got head_dim={self.head_dim}." ) def forward( self, x: torch.Tensor, key_padding_mask: Optional[torch.Tensor] = None, ) -> torch.Tensor: # x: [B, T, C], key_padding_mask: [B, T] with True = PAD b, t, c = x.shape q, k, v = self.qkv(x).chunk(3, dim=-1) q = q.view(b, t, self.num_heads, self.head_dim).transpose(1, 2) k = k.view(b, t, self.num_heads, self.head_dim).transpose(1, 2) v = v.view(b, t, self.num_heads, self.head_dim).transpose(1, 2) if self.use_rope: cos, sin = self.rope.get_cos_sin(t, x.device, x.dtype) q = apply_rope(q, cos, sin) k = apply_rope(k, cos, sin) attn_scores = torch.matmul(q, k.transpose(-2, -1)) * self.scale if key_padding_mask is not None: pad = key_padding_mask.unsqueeze(1).unsqueeze(2) # [B,1,1,T] attn_scores = attn_scores.masked_fill( pad, torch.finfo(attn_scores.dtype).min ) attn = F.softmax(attn_scores, dim=-1) attn = self.attn_dropout(attn) out = torch.matmul(attn, v) # [B,H,T,D] out = out.transpose(1, 2).contiguous().view(b, t, c) return self.out_proj(out) class HRMReasoningBlock(nn.Module): def __init__( self, dim: int, num_heads: int, expansion: float, style_dim: int, attn_dropout: float = 0.0, resid_dropout: float = 0.0, use_rope: bool = True, rope_base: float = 10000.0, rms_eps: float = 1e-5, ): super().__init__() self.attn = MultiheadSelfAttentionRoPE( dim, num_heads, attn_dropout, use_rope, rope_base ) self.ff = SwiGLUFFN(dim, expansion, resid_dropout) self.resid_dropout = nn.Dropout(resid_dropout) self.ada1 = AdaRMSNorm(style_dim, dim, eps=rms_eps) self.ada2 = AdaRMSNorm(style_dim, dim, eps=rms_eps) def forward( self, hidden_states: torch.Tensor, style: torch.Tensor, key_padding_mask: Optional[torch.Tensor] = None, ) -> torch.Tensor: attn_out = self.attn(hidden_states, key_padding_mask=key_padding_mask) hidden_states = self.ada1( hidden_states + self.resid_dropout(attn_out), style ) mlp_out = self.ff(hidden_states) hidden_states = self.ada2( hidden_states + self.resid_dropout(mlp_out), style ) return hidden_states class HRMReasoningModule(nn.Module): def __init__(self, layers: nn.ModuleList): super().__init__() self.layers = layers def forward( self, hidden_states: torch.Tensor, input_injection: torch.Tensor, style: torch.Tensor, key_padding_mask: Optional[torch.Tensor] = None, ) -> torch.Tensor: hidden_states = hidden_states + input_injection for layer in self.layers: hidden_states = layer( hidden_states, style, key_padding_mask=key_padding_mask ) return hidden_states # ----------------------------------------------------------------------------- # HRM-based StyleTransformer1d (same external API as your current one) # ----------------------------------------------------------------------------- class StyleTransformer1d(nn.Module): """ HRM replacement for your StyleTransformer1d. - Works on tokens: concat([x, embedding]) => model_dim = channels + cond_dim - Uses HRM H/L cycles instead of stacked transformer blocks - Style-conditioning via AdaRMSNorm in every block - Time/features mapping added as a bias/injection """ def __init__( self, num_layers: int, channels: int, num_heads: int, head_features: int, # kept for signature compatibility; ignored multiplier: int, use_context_time: bool = True, use_rel_pos: bool = False, # ignored context_features_multiplier: int = 1, # ignored rel_pos_num_buckets: Optional[int] = None, # ignored rel_pos_max_distance: Optional[int] = None, # ignored context_features: Optional[int] = None, # style dim context_embedding_features: Optional[int] = None, # cond dim embedding_max_length: int = 512, H_cycles: int = 1, L_cycles: int = 2, attn_dropout: float = 0.0, resid_dropout: float = 0.0, use_rope: bool = True, rope_base: float = 10000.0, ): super().__init__() assert exists(context_embedding_features), "context_embedding_features is required" assert exists(context_features), "context_features (style dim) is required" self.channels = channels self.cond_dim = context_embedding_features self.style_dim = context_features self.model_dim = channels + context_embedding_features self.use_context_time = use_context_time self.use_context_features = exists(context_features) self.H_cycles = H_cycles self.L_cycles = L_cycles # Optional positional embedding added to conditioning sequence self.fixed_embedding = FixedEmbedding( max_length=embedding_max_length, features=context_embedding_features, ) # time/features mapping -> model_dim (same as concat(x, embedding)) if use_context_time or self.use_context_features: self.to_mapping = nn.Sequential( nn.Linear(self.model_dim, self.model_dim), nn.GELU(), nn.Linear(self.model_dim, self.model_dim), nn.GELU(), ) if use_context_time: self.to_time = nn.Sequential( TimePositionalEmbedding(dim=channels, out_features=self.model_dim), nn.GELU(), ) if self.use_context_features: self.to_features = nn.Sequential( nn.Linear(context_features, self.model_dim), nn.GELU(), ) # HRM states + step embeddings self.H_init = nn.Parameter(torch.zeros(1, 1, self.model_dim)) self.L_init = nn.Parameter(torch.zeros(1, 1, self.model_dim)) nn.init.trunc_normal_(self.H_init, std=0.02) nn.init.trunc_normal_(self.L_init, std=0.02) self.h_step_emb = nn.Embedding(H_cycles, self.model_dim) self.l_step_emb = nn.Embedding(L_cycles, self.model_dim) layers_L = nn.ModuleList( [ HRMReasoningBlock( dim=self.model_dim, num_heads=num_heads, expansion=float(multiplier), style_dim=self.style_dim, attn_dropout=attn_dropout, resid_dropout=resid_dropout, use_rope=use_rope, rope_base=rope_base, ) for _ in range(num_layers) ] ) layers_H = nn.ModuleList( [ HRMReasoningBlock( dim=self.model_dim, num_heads=num_heads, expansion=float(multiplier), style_dim=self.style_dim, attn_dropout=attn_dropout, resid_dropout=resid_dropout, use_rope=use_rope, rope_base=rope_base, ) for _ in range(num_layers) ] ) self.L_level = HRMReasoningModule(layers_L) self.H_level = HRMReasoningModule(layers_H) # project back to `channels` self.to_out = nn.Sequential( Rearrange("b t c -> b c t"), nn.Conv1d(self.model_dim, channels, kernel_size=1), ) def get_mapping( self, time: Optional[torch.Tensor] = None, features: Optional[torch.Tensor] = None, ) -> Optional[torch.Tensor]: items = [] if self.use_context_time: assert exists(time), "use_context_time=True but time was not provided" items.append(self.to_time(time)) if self.use_context_features: assert exists(features), "context_features exists but features were not provided" items.append(self.to_features(features)) if len(items) == 0: return None mapping = torch.stack(items, dim=0).sum(dim=0) # [B, model_dim] return self.to_mapping(mapping) def run( self, x: torch.Tensor, # [B, T, channels] time: torch.Tensor, # [B]a embedding: torch.Tensor, # [B, T, cond_dim] features: torch.Tensor, # [B, style_dim] key_padding_mask: Optional[torch.Tensor] = None, # [B, T] True=PAD ) -> torch.Tensor: # add positional embedding to cond (this was computed but unused in your old code) embedding = embedding + self.fixed_embedding(embedding).to(dtype=embedding.dtype) # concat token-wise x_cat = torch.cat([x, embedding], dim=-1) # [B, T, model_dim] mapping = self.get_mapping(time=time, features=features) if mapping is not None: x_cat = x_cat + mapping.unsqueeze(1) b, t, _ = x_cat.shape z_H = self.H_init.to(dtype=x_cat.dtype).expand(b, t, -1) z_L = self.L_init.to(dtype=x_cat.dtype).expand(b, t, -1) for h in range(self.H_cycles): h_emb = self.h_step_emb.weight[h].view(1, 1, -1).to(dtype=x_cat.dtype) z_H_curr = z_H + h_emb for l in range(self.L_cycles): l_emb = self.l_step_emb.weight[l].view(1, 1, -1).to(dtype=x_cat.dtype) inj = z_H_curr + x_cat + l_emb z_L = self.L_level( z_L, input_injection=inj, style=features, key_padding_mask=key_padding_mask, ) if key_padding_mask is not None: z_L = z_L.masked_fill(key_padding_mask.unsqueeze(-1), 0.0) z_H = self.H_level( z_H, input_injection=z_L, style=features, key_padding_mask=key_padding_mask, ) if key_padding_mask is not None: z_H = z_H.masked_fill(key_padding_mask.unsqueeze(-1), 0.0) out = self.to_out(z_H).transpose(1, 2) # [B, T, channels] return out def forward( self, x: torch.Tensor, time: torch.Tensor, embedding_mask_proba: float = 0.0, embedding: Optional[torch.Tensor] = None, features: Optional[torch.Tensor] = None, embedding_scale: float = 1.0, key_padding_mask: Optional[torch.Tensor] = None, ) -> torch.Tensor: assert exists(embedding), "embedding is required" assert exists(features), "features(style) is required" b = embedding.shape[0] device = embedding.device if embedding_mask_proba > 0.0: # Shape [B, 1, 1] for broadcasting batch_mask = rand_bool((b, 1, 1), embedding_mask_proba, device=device) # Mask the embedding [B, T, D] embedding = torch.where(batch_mask, torch.zeros_like(embedding), embedding) # Mask the features [B, S] -> squeeze mask to [B, 1] for broadcasting features_mask = batch_mask.squeeze(-1) # [B, 1] features = torch.where(features_mask, torch.zeros_like(features), features) if embedding_scale != 1.0: out = self.run( x=x, time=time, embedding=embedding, features=features, key_padding_mask=key_padding_mask, ) null_embedding = torch.zeros_like(embedding) out_null = self.run( x=x, time=time, embedding=null_embedding, features=features, key_padding_mask=key_padding_mask, ) return out_null + (out - out_null) * embedding_scale return self.run( x=x, time=time, embedding=embedding, features=features, key_padding_mask=key_padding_mask, )