"""SWaG backbone adapted from the DiT architecture for 1-D waveforms. The multi-scale waveform frontend and decoder follow the SWaG research code. The transformer conditioning path contains diffusion timesteps only. """ from __future__ import annotations import math import numpy as np import torch import torch.nn as nn import torch.nn.functional as F from timm.models.vision_transformer import Attention, Mlp def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1) def sinusoidal_position_embedding(length: int, dimension: int) -> torch.Tensor: if dimension % 2: raise ValueError("Positional embedding dimension must be even") positions = np.arange(length, dtype=np.float32)[:, None] frequencies = np.exp(-math.log(10_000) * np.arange(dimension // 2) / (dimension // 2)) embedding = np.concatenate([np.sin(positions * frequencies), np.cos(positions * frequencies)], axis=1) return torch.from_numpy(embedding.astype(np.float32)).unsqueeze(0) class TimestepEmbedder(nn.Module): def __init__(self, hidden_size: int, frequency_size: int = 256): super().__init__() self.frequency_size = frequency_size self.mlp = nn.Sequential( nn.Linear(frequency_size, hidden_size), nn.SiLU(), nn.Linear(hidden_size, hidden_size), ) @staticmethod def frequency_embedding(t: torch.Tensor, dimension: int, max_period: int = 10_000) -> torch.Tensor: half = dimension // 2 frequencies = torch.exp( -math.log(max_period) * torch.arange(half, dtype=torch.float32, device=t.device) / half ) angles = t[:, None].float() * frequencies[None] embedding = torch.cat([torch.cos(angles), torch.sin(angles)], dim=-1) if dimension % 2: embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) return embedding def forward(self, timesteps: torch.Tensor) -> torch.Tensor: return self.mlp(self.frequency_embedding(timesteps, self.frequency_size)) class ContinuousConditionEmbedder(nn.Module): """Embed one continuous scalar and provide a learned null value for CFG.""" def __init__(self, hidden_size: int, frequency_size: int = 256): super().__init__() self.frequency_size = frequency_size self.mlp = nn.Sequential( nn.Linear(frequency_size, hidden_size), nn.SiLU(), nn.Linear(hidden_size, hidden_size), ) self.null_embedding = nn.Parameter(torch.zeros(hidden_size)) def forward(self, values: torch.Tensor, drop_mask: torch.Tensor) -> torch.Tensor: embedded = self.mlp(TimestepEmbedder.frequency_embedding(values, self.frequency_size)) return torch.where(drop_mask[:, None], self.null_embedding[None].to(embedded.dtype), embedded) class ConvBlock(nn.Module): def __init__(self, in_channels: int, out_channels: int, kernel_size: int): super().__init__() self.block = nn.Sequential( nn.Conv1d(in_channels, out_channels, kernel_size, padding=kernel_size // 2), nn.GroupNorm(1, out_channels), nn.GELU(), ) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.block(x) class TokenBranch(nn.Module): def __init__(self, in_channels: int, branch_channels: int, hidden_size: int, tokens: int, kernel: int): super().__init__() self.tokens = tokens self.encoder = nn.Sequential( ConvBlock(in_channels, branch_channels, kernel), ConvBlock(branch_channels, branch_channels, kernel), ) self.pool = nn.AdaptiveAvgPool1d(tokens) self.projection = nn.Conv1d(branch_channels, hidden_size, 1) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.projection(self.pool(self.encoder(x))).transpose(1, 2) class MultiScaleEncoder(nn.Module): def __init__( self, in_channels: int, hidden_size: int, stem_channels: int, small_channels: int, mid_channels: int, large_channels: int, small_tokens: int, mid_tokens: int, large_tokens: int, ): super().__init__() self.stem = nn.Sequential( ConvBlock(in_channels, stem_channels, 9), ConvBlock(stem_channels, stem_channels, 7), ) self.small = TokenBranch(stem_channels, small_channels, hidden_size, small_tokens, 7) self.mid = TokenBranch(stem_channels, mid_channels, hidden_size, mid_tokens, 9) self.large = TokenBranch(stem_channels, large_channels, hidden_size, large_tokens, 15) def forward(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: features = self.stem(x) return self.small(features), self.mid(features), self.large(features) class TokenPreprocessor(nn.Module): def __init__(self, hidden_size: int): super().__init__() self.norms = nn.ModuleList([nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) for _ in range(3)]) self.modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 6 * hidden_size)) def forward(self, small: torch.Tensor, mid: torch.Tensor, large: torch.Tensor, t: torch.Tensor): shifts = self.modulation(t).chunk(6, dim=1) return ( modulate(self.norms[0](small), shifts[0], shifts[1]), modulate(self.norms[1](mid), shifts[2], shifts[3]), modulate(self.norms[2](large), shifts[4], shifts[5]), ) class MultiScaleDiTBlock(nn.Module): def __init__(self, hidden_size: int, num_heads: int, mlp_ratio: float): super().__init__() self.norm_self = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) self.self_attention = Attention(hidden_size, num_heads=num_heads, qkv_bias=True) self.norm_mid_query = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) self.norm_mid_context = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) self.mid_attention = nn.MultiheadAttention(hidden_size, num_heads, batch_first=True) self.norm_large_query = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) self.norm_large_context = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) self.large_attention = nn.MultiheadAttention(hidden_size, num_heads, batch_first=True) self.norm_mlp = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6) self.mlp = Mlp( in_features=hidden_size, hidden_features=int(hidden_size * mlp_ratio), act_layer=lambda: nn.GELU(approximate="tanh"), drop=0, ) self.modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden_size, 16 * hidden_size)) def forward(self, small: torch.Tensor, mid: torch.Tensor, large: torch.Tensor, t: torch.Tensor) -> torch.Tensor: values = self.modulation(t).chunk(16, dim=1) small = small + values[2].unsqueeze(1) * self.self_attention( modulate(self.norm_self(small), values[0], values[1]) ) query = modulate(self.norm_mid_query(small), values[3], values[4]) context = modulate(self.norm_mid_context(mid), values[5], values[6]) update, _ = self.mid_attention(query, context, context, need_weights=False) small = small + values[7].unsqueeze(1) * update query = modulate(self.norm_large_query(small), values[8], values[9]) context = modulate(self.norm_large_context(large), values[10], values[11]) update, _ = self.large_attention(query, context, context, need_weights=False) small = small + values[12].unsqueeze(1) * update return small + values[15].unsqueeze(1) * self.mlp( modulate(self.norm_mlp(small), values[13], values[14]) ) class LoRALinearDelta(nn.Module): def __init__(self, features: int, rank: int, alpha: float, dropout: float): super().__init__() self.scale = float(alpha) / int(rank) self.dropout = nn.Dropout(float(dropout)) if dropout > 0 else nn.Identity() self.down = nn.Linear(features, rank, bias=False) self.up = nn.Linear(rank, features, bias=False) nn.init.kaiming_uniform_(self.down.weight, a=math.sqrt(5)) nn.init.zeros_(self.up.weight) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.up(self.down(self.dropout(x))) * self.scale class LoRAMultiheadAttention(nn.MultiheadAttention): """MultiheadAttention with additive LoRA while retaining base state keys.""" def __init__(self, embed_dim: int, num_heads: int, rank: int, alpha: float, dropout: float): super().__init__(embed_dim, num_heads, batch_first=True) self.lora_q = LoRALinearDelta(embed_dim, rank, alpha, dropout) self.lora_k = LoRALinearDelta(embed_dim, rank, alpha, dropout) self.lora_v = LoRALinearDelta(embed_dim, rank, alpha, dropout) self.lora_out = LoRALinearDelta(embed_dim, rank, alpha, dropout) def forward(self, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, **kwargs): dimension = self.embed_dim q_weight, k_weight, v_weight = self.in_proj_weight.chunk(3, dim=0) q_bias, k_bias, v_bias = self.in_proj_bias.chunk(3, dim=0) q = F.linear(query, q_weight, q_bias) + self.lora_q(query) k = F.linear(key, k_weight, k_bias) + self.lora_k(key) v = F.linear(value, v_weight, v_bias) + self.lora_v(value) batch, query_length, _ = q.shape key_length = k.shape[1] head_dimension = dimension // self.num_heads q = q.view(batch, query_length, self.num_heads, head_dimension).transpose(1, 2) k = k.view(batch, key_length, self.num_heads, head_dimension).transpose(1, 2) v = v.view(batch, key_length, self.num_heads, head_dimension).transpose(1, 2) attended = F.scaled_dot_product_attention(q, k, v, dropout_p=0.0) attended = attended.transpose(1, 2).contiguous().view(batch, query_length, dimension) return self.out_proj(attended) + self.lora_out(attended), None class SWaG(nn.Module): """SWaG with the fixed multi-scale cross-attention waveform backbone.""" def __init__( self, in_channels: int = 3, length: int = 6000, hidden_size: int = 768, depth: int = 12, num_heads: int = 12, mlp_ratio: float = 4.0, frequency_embedding_size: int = 256, small_token_count: int = 1024, mid_token_count: int = 256, large_token_count: int = 32, stem_feature_channels: int = 64, small_encoder_channels: int = 128, mid_encoder_channels: int = 128, large_encoder_channels: int = 128, decoder_channels: int = 128, skip_channels: int = 64, fusion_channels: int = 128, use_high_res_skip: bool = True, learn_sigma: bool = True, ): super().__init__() if hidden_size % num_heads: raise ValueError("hidden_size must be divisible by num_heads") self.in_channels = in_channels self.length = length self.learn_sigma = learn_sigma self.out_channels = in_channels * 2 if learn_sigma else in_channels self.use_high_res_skip = use_high_res_skip self.encoder = MultiScaleEncoder( in_channels, hidden_size, stem_feature_channels, small_encoder_channels, mid_encoder_channels, large_encoder_channels, small_token_count, mid_token_count, large_token_count, ) self.timestep_embedder = TimestepEmbedder(hidden_size, frequency_embedding_size) self.register_buffer("small_position", sinusoidal_position_embedding(small_token_count, hidden_size), persistent=True) self.register_buffer("mid_position", sinusoidal_position_embedding(mid_token_count, hidden_size), persistent=True) self.register_buffer("large_position", sinusoidal_position_embedding(large_token_count, hidden_size), persistent=True) self.token_preprocessor = TokenPreprocessor(hidden_size) self.blocks = nn.ModuleList([MultiScaleDiTBlock(hidden_size, num_heads, mlp_ratio) for _ in range(depth)]) self.decoder_projection = nn.Conv1d(hidden_size, decoder_channels, 1) self.decoder_refine = ConvBlock(decoder_channels, decoder_channels, 5) self.skip_branch = ( nn.Sequential(ConvBlock(in_channels, skip_channels, 7), ConvBlock(skip_channels, skip_channels, 5)) if use_high_res_skip else None ) fusion_input = decoder_channels + (skip_channels if use_high_res_skip else 0) self.fusion = ConvBlock(fusion_input, fusion_channels, 5) self.output_head = nn.Conv1d(fusion_channels, self.out_channels, 1) self.initialize_weights() def initialize_weights(self) -> None: def initialize(module: nn.Module) -> None: if isinstance(module, nn.Linear): nn.init.xavier_uniform_(module.weight) if module.bias is not None: nn.init.zeros_(module.bias) self.apply(initialize) nn.init.normal_(self.timestep_embedder.mlp[0].weight, std=0.02) nn.init.normal_(self.timestep_embedder.mlp[2].weight, std=0.02) nn.init.zeros_(self.token_preprocessor.modulation[-1].weight) nn.init.zeros_(self.token_preprocessor.modulation[-1].bias) for block in self.blocks: nn.init.zeros_(block.modulation[-1].weight) nn.init.zeros_(block.modulation[-1].bias) nn.init.zeros_(self.output_head.weight) nn.init.zeros_(self.output_head.bias) def forward(self, x: torch.Tensor, timesteps: torch.Tensor) -> torch.Tensor: if x.ndim != 3 or x.shape[1:] != (self.in_channels, self.length): raise ValueError(f"Expected x shape [N, {self.in_channels}, {self.length}], got {tuple(x.shape)}") if timesteps.ndim != 1 or timesteps.shape[0] != x.shape[0]: raise ValueError(f"Expected timesteps shape [{x.shape[0]}], got {tuple(timesteps.shape)}") timestep_embedding = self.timestep_embedder(timesteps) small, mid, large = self.encoder(x) small = small + self.small_position.to(dtype=small.dtype) mid = mid + self.mid_position.to(dtype=mid.dtype) large = large + self.large_position.to(dtype=large.dtype) small, mid, large = self.token_preprocessor(small, mid, large, timestep_embedding) for block in self.blocks: small = block(small, mid, large, timestep_embedding) decoded = self.decoder_projection(small.transpose(1, 2)) decoded = F.interpolate(decoded, size=self.length, mode="linear", align_corners=False) decoded = self.decoder_refine(decoded) if self.skip_branch is not None: decoded = torch.cat([decoded, self.skip_branch(x)], dim=1) return self.output_head(self.fusion(decoded)) class EmptyConditionSWaG(SWaG): """No-skip SWaG with eight reserved continuous condition slots. The slots are passed as normalized scalar values. Their projection is zero-initialized, so an all-zero condition vector leaves the base model behavior unchanged while preserving a migration interface. """ def __init__( self, condition_slot_count: int = 8, condition_normalization_length: float = 6000.0, hidden_size: int = 768, frequency_embedding_size: int = 256, **kwargs, ): super().__init__( hidden_size=hidden_size, frequency_embedding_size=frequency_embedding_size, **kwargs, ) if int(condition_slot_count) != 8: raise ValueError("EmptyConditionSWaG requires exactly 8 condition slots") self.condition_slot_count = 8 self.condition_normalization_length = float(condition_normalization_length) hidden = int(hidden_size) frequency_size = int(frequency_embedding_size) self.condition_embedders = nn.ModuleList( [ContinuousConditionEmbedder(hidden, frequency_size) for _ in range(self.condition_slot_count)] ) self.condition_projection = nn.Sequential(nn.SiLU(), nn.Linear(hidden, hidden)) nn.init.zeros_(self.condition_projection[-1].weight) nn.init.zeros_(self.condition_projection[-1].bias) def forward( self, x: torch.Tensor, timesteps: torch.Tensor, conditions: torch.Tensor | None = None, ) -> torch.Tensor: if conditions is None: conditions = torch.zeros( x.shape[0], self.condition_slot_count, device=x.device, dtype=torch.float32 ) if conditions.ndim != 2 or conditions.shape != (x.shape[0], self.condition_slot_count): raise ValueError( f"Expected conditions shape [{x.shape[0]}, {self.condition_slot_count}], " f"got {tuple(conditions.shape)}" ) if not torch.isfinite(conditions).all(): raise ValueError("Conditions contain non-finite values") normalized = conditions.float() / self.condition_normalization_length drop_mask = torch.zeros(x.shape[0], dtype=torch.bool, device=x.device) encoded = sum( embedder(normalized[:, index], drop_mask) for index, embedder in enumerate(self.condition_embedders) ) / math.sqrt(self.condition_slot_count) timestep_embedding = self.timestep_embedder(timesteps) + self.condition_projection(encoded) small, mid, large = self.encoder(x) small = small + self.small_position.to(dtype=small.dtype) mid = mid + self.mid_position.to(dtype=mid.dtype) large = large + self.large_position.to(dtype=large.dtype) small, mid, large = self.token_preprocessor(small, mid, large, timestep_embedding) for block in self.blocks: small = block(small, mid, large, timestep_embedding) decoded = self.decoder_projection(small.transpose(1, 2)) decoded = F.interpolate(decoded, size=self.length, mode="linear", align_corners=False) decoded = self.decoder_refine(decoded) return self.output_head(self.fusion(decoded)) class ConditionalSWaG(SWaG): """SWaG conditioned jointly on P- and S-arrival sample indices.""" def __init__(self, condition_dropout_prob: float = 0.1, **kwargs): super().__init__(**kwargs) if not 0.0 <= condition_dropout_prob <= 1.0: raise ValueError("condition_dropout_prob must be in [0, 1]") hidden_size = self.timestep_embedder.mlp[-1].out_features frequency_size = self.timestep_embedder.frequency_size self.condition_dropout_prob = float(condition_dropout_prob) self.p_embedder = ContinuousConditionEmbedder(hidden_size, frequency_size) self.s_embedder = ContinuousConditionEmbedder(hidden_size, frequency_size) for embedder in (self.p_embedder, self.s_embedder): nn.init.normal_(embedder.mlp[0].weight, std=0.02) nn.init.normal_(embedder.mlp[2].weight, std=0.02) def _drop_mask(self, labels: torch.Tensor, force_drop_mask: torch.Tensor | None) -> torch.Tensor: if force_drop_mask is not None: mask = force_drop_mask.to(device=labels.device, dtype=torch.bool) if mask.shape != (labels.shape[0],): raise ValueError(f"Expected force_drop_mask shape [{labels.shape[0]}], got {tuple(mask.shape)}") return mask if self.training and self.condition_dropout_prob > 0: return torch.rand(labels.shape[0], device=labels.device) < self.condition_dropout_prob return torch.zeros(labels.shape[0], dtype=torch.bool, device=labels.device) def forward( self, x: torch.Tensor, timesteps: torch.Tensor, labels: torch.Tensor, force_drop_mask: torch.Tensor | None = None, ) -> torch.Tensor: if x.ndim != 3 or x.shape[1:] != (self.in_channels, self.length): raise ValueError(f"Expected x shape [N, {self.in_channels}, {self.length}], got {tuple(x.shape)}") if timesteps.ndim != 1 or timesteps.shape[0] != x.shape[0]: raise ValueError(f"Expected timesteps shape [{x.shape[0]}], got {tuple(timesteps.shape)}") if labels.shape != (x.shape[0], 2): raise ValueError(f"Expected P/S labels shape [{x.shape[0]}, 2], got {tuple(labels.shape)}") if not torch.isfinite(labels).all(): raise ValueError("P/S labels contain non-finite values") drop_mask = self._drop_mask(labels, force_drop_mask) condition = (self.p_embedder(labels[:, 0], drop_mask) + self.s_embedder(labels[:, 1], drop_mask)) / math.sqrt(2.0) conditioning = self.timestep_embedder(timesteps) + condition small, mid, large = self.encoder(x) small = small + self.small_position.to(dtype=small.dtype) mid = mid + self.mid_position.to(dtype=mid.dtype) large = large + self.large_position.to(dtype=large.dtype) small, mid, large = self.token_preprocessor(small, mid, large, conditioning) for block in self.blocks: small = block(small, mid, large, conditioning) decoded = self.decoder_projection(small.transpose(1, 2)) decoded = F.interpolate(decoded, size=self.length, mode="linear", align_corners=False) decoded = self.decoder_refine(decoded) if self.skip_branch is not None: decoded = torch.cat([decoded, self.skip_branch(x)], dim=1) return self.output_head(self.fusion(decoded)) def forward_with_cfg( self, x: torch.Tensor, timesteps: torch.Tensor, labels: torch.Tensor, cfg_scale: float ) -> torch.Tensor: """Run conditional and null passes and apply CFG to epsilon channels only.""" keep = torch.zeros(x.shape[0], dtype=torch.bool, device=x.device) drop = torch.ones(x.shape[0], dtype=torch.bool, device=x.device) conditional = self(x, timesteps, labels, force_drop_mask=keep) unconditional = self(x, timesteps, labels, force_drop_mask=drop) eps_c, rest_c = conditional[:, : self.in_channels], conditional[:, self.in_channels :] eps_u = unconditional[:, : self.in_channels] guided_eps = eps_u + float(cfg_scale) * (eps_c - eps_u) return torch.cat([guided_eps, rest_c], dim=1) class ConditionalLoRASWaG(ConditionalSWaG): """P/S-conditioned SWaG with normalized arrivals and cross-attention LoRA.""" def __init__( self, condition_normalization_length: float = 6000.0, cross_attention_lora_rank: int = 8, cross_attention_lora_alpha: float = 16.0, cross_attention_lora_dropout: float = 0.05, **kwargs, ): super().__init__(**kwargs) self.condition_normalization_length = float(condition_normalization_length) for block in self.blocks: for name in ("mid_attention", "large_attention"): base = getattr(block, name) adapted = LoRAMultiheadAttention( base.embed_dim, base.num_heads, cross_attention_lora_rank, cross_attention_lora_alpha, cross_attention_lora_dropout, ) adapted.in_proj_weight.data.copy_(base.in_proj_weight.data) adapted.in_proj_bias.data.copy_(base.in_proj_bias.data) adapted.out_proj.load_state_dict(base.out_proj.state_dict()) setattr(block, name, adapted) def forward(self, x, timesteps, labels, force_drop_mask=None): normalized = labels / self.condition_normalization_length return super().forward(x, timesteps, normalized, force_drop_mask=force_drop_mask) class DualConditionLoRASWaG(ConditionalLoRASWaG): """LoRA SWaG with separate timestep and P/S adaLN modulation chains.""" def __init__(self, **kwargs): super().__init__(**kwargs) hidden = self.timestep_embedder.mlp[-1].out_features self.token_preprocessor.condition_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden, 6 * hidden)) nn.init.zeros_(self.token_preprocessor.condition_modulation[-1].weight) nn.init.zeros_(self.token_preprocessor.condition_modulation[-1].bias) for block in self.blocks: block.condition_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden, 16 * hidden)) nn.init.zeros_(block.condition_modulation[-1].weight) nn.init.zeros_(block.condition_modulation[-1].bias) def forward(self, x, timesteps, labels, force_drop_mask=None): if x.ndim != 3 or x.shape[1:] != (self.in_channels, self.length): raise ValueError(f"Expected x shape [N, {self.in_channels}, {self.length}], got {tuple(x.shape)}") if labels.shape != (x.shape[0], 2): raise ValueError(f"Expected labels shape [{x.shape[0]}, 2], got {tuple(labels.shape)}") normalized = labels / self.condition_normalization_length drop_mask = self._drop_mask(normalized, force_drop_mask) condition = (self.p_embedder(normalized[:, 0], drop_mask) + self.s_embedder(normalized[:, 1], drop_mask)) / math.sqrt(2.0) timestep = self.timestep_embedder(timesteps) small, mid, large = self.encoder(x) small = small + self.small_position.to(dtype=small.dtype) mid = mid + self.mid_position.to(dtype=mid.dtype) large = large + self.large_position.to(dtype=large.dtype) time_values = self.token_preprocessor.modulation(timestep).chunk(6, dim=1) cond_values = self.token_preprocessor.condition_modulation(condition).chunk(6, dim=1) small, mid, large = ( modulate(self.token_preprocessor.norms[0](small), time_values[0] + cond_values[0], time_values[1] + cond_values[1]), modulate(self.token_preprocessor.norms[1](mid), time_values[2] + cond_values[2], time_values[3] + cond_values[3]), modulate(self.token_preprocessor.norms[2](large), time_values[4] + cond_values[4], time_values[5] + cond_values[5]), ) for block in self.blocks: time_values = block.modulation(timestep).chunk(16, dim=1) cond_values = block.condition_modulation(condition).chunk(16, dim=1) values = [time_values[i] + cond_values[i] for i in range(16)] small = small + values[2].unsqueeze(1) * block.self_attention(modulate(block.norm_self(small), values[0], values[1])) query = modulate(block.norm_mid_query(small), values[3], values[4]) context = modulate(block.norm_mid_context(mid), values[5], values[6]) update, _ = block.mid_attention(query, context, context, need_weights=False) small = small + values[7].unsqueeze(1) * update query = modulate(block.norm_large_query(small), values[8], values[9]) context = modulate(block.norm_large_context(large), values[10], values[11]) update, _ = block.large_attention(query, context, context, need_weights=False) small = small + values[12].unsqueeze(1) * update small = small + values[15].unsqueeze(1) * block.mlp(modulate(block.norm_mlp(small), values[13], values[14])) decoded = self.decoder_projection(small.transpose(1, 2)) decoded = F.interpolate(decoded, size=self.length, mode="linear", align_corners=False) decoded = self.decoder_refine(decoded) if self.skip_branch is not None: decoded = torch.cat([decoded, self.skip_branch(x)], dim=1) return self.output_head(self.fusion(decoded)) class AdaLNSumSWaG(ConditionalSWaG): """Frozen-backbone SWaG with independent P/S modulation parameters added to timestep parameters.""" def __init__(self, condition_normalization_length: float = 6000.0, **kwargs): super().__init__(**kwargs) self.condition_normalization_length = float(condition_normalization_length) hidden = self.timestep_embedder.mlp[-1].out_features self.token_preprocessor.condition_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden, 6 * hidden)) nn.init.zeros_(self.token_preprocessor.condition_modulation[-1].weight) nn.init.zeros_(self.token_preprocessor.condition_modulation[-1].bias) for block in self.blocks: block.condition_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden, 16 * hidden)) nn.init.zeros_(block.condition_modulation[-1].weight) nn.init.zeros_(block.condition_modulation[-1].bias) def _embeddings(self, timesteps, labels, force_drop_mask): normalized = labels / self.condition_normalization_length drop = self._drop_mask(normalized, force_drop_mask) condition = (self.p_embedder(normalized[:, 0], drop) + self.s_embedder(normalized[:, 1], drop)) / math.sqrt(2.0) return self.timestep_embedder(timesteps), condition def _decode(self, small, x): decoded = self.decoder_projection(small.transpose(1, 2)) decoded = self.decoder_refine(F.interpolate(decoded, size=self.length, mode="linear", align_corners=False)) if self.skip_branch is not None: decoded = torch.cat([decoded, self.skip_branch(x)], dim=1) return self.output_head(self.fusion(decoded)) def forward(self, x, timesteps, labels, force_drop_mask=None): timestep, condition = self._embeddings(timesteps, labels, force_drop_mask) small, mid, large = self.encoder(x) small, mid, large = small + self.small_position, mid + self.mid_position, large + self.large_position tv = self.token_preprocessor.modulation(timestep).chunk(6, 1) cv = self.token_preprocessor.condition_modulation(condition).chunk(6, 1) small, mid, large = tuple( modulate(self.token_preprocessor.norms[i](token), tv[2*i] + cv[2*i], tv[2*i+1] + cv[2*i+1]) for i, token in enumerate((small, mid, large)) ) for block in self.blocks: t = block.modulation(timestep).chunk(16, 1); c = block.condition_modulation(condition).chunk(16, 1) v = [t[i] + c[i] for i in range(16)] small = small + v[2].unsqueeze(1) * block.self_attention(modulate(block.norm_self(small), v[0], v[1])) q = modulate(block.norm_mid_query(small), v[3], v[4]); ctx = modulate(block.norm_mid_context(mid), v[5], v[6]) update, _ = block.mid_attention(q, ctx, ctx, need_weights=False); small = small + v[7].unsqueeze(1) * update q = modulate(block.norm_large_query(small), v[8], v[9]); ctx = modulate(block.norm_large_context(large), v[10], v[11]) update, _ = block.large_attention(q, ctx, ctx, need_weights=False); small = small + v[12].unsqueeze(1) * update small = small + v[15].unsqueeze(1) * block.mlp(modulate(block.norm_mlp(small), v[13], v[14])) return self._decode(small, x) class AdaLNResidualSWaG(AdaLNSumSWaG): """SWaG with timestep and P/S modulation applied as separate residual updates.""" def __init__(self, **kwargs): super().__init__(**kwargs) hidden = self.timestep_embedder.mlp[-1].out_features self.token_preprocessor.condition_modulation = nn.Sequential(nn.SiLU(), nn.Linear(hidden, 9 * hidden)) nn.init.zeros_(self.token_preprocessor.condition_modulation[-1].weight) nn.init.zeros_(self.token_preprocessor.condition_modulation[-1].bias) def forward(self, x, timesteps, labels, force_drop_mask=None): timestep, condition = self._embeddings(timesteps, labels, force_drop_mask) small, mid, large = self.encoder(x) small, mid, large = small + self.small_position, mid + self.mid_position, large + self.large_position original = (small, mid, large) tv = self.token_preprocessor.modulation(timestep).chunk(6, 1) cv = self.token_preprocessor.condition_modulation(condition).chunk(9, 1) tokens = [] for i, token in enumerate(original): norm = self.token_preprocessor.norms[i](token) time_token = modulate(norm, tv[2*i], tv[2*i+1]) cond_token = modulate(norm, cv[3*i], cv[3*i+1]) tokens.append(time_token + cv[3*i+2].unsqueeze(1) * cond_token) small, mid, large = tokens for block in self.blocks: t = block.modulation(timestep).chunk(16, 1); c = block.condition_modulation(condition).chunk(16, 1) norm = block.norm_self(small) small = small + t[2].unsqueeze(1) * block.self_attention(modulate(norm, t[0], t[1])) small = small + c[2].unsqueeze(1) * block.self_attention(modulate(norm, c[0], c[1])) qt = modulate(block.norm_mid_query(small), t[3], t[4]); ct = modulate(block.norm_mid_context(mid), t[5], t[6]) qc = modulate(block.norm_mid_query(small), c[3], c[4]); cc = modulate(block.norm_mid_context(mid), c[5], c[6]) u, _ = block.mid_attention(qt, ct, ct, need_weights=False); small = small + t[7].unsqueeze(1) * u u, _ = block.mid_attention(qc, cc, cc, need_weights=False); small = small + c[7].unsqueeze(1) * u qt = modulate(block.norm_large_query(small), t[8], t[9]); ct = modulate(block.norm_large_context(large), t[10], t[11]) qc = modulate(block.norm_large_query(small), c[8], c[9]); cc = modulate(block.norm_large_context(large), c[10], c[11]) u, _ = block.large_attention(qt, ct, ct, need_weights=False); small = small + t[12].unsqueeze(1) * u u, _ = block.large_attention(qc, cc, cc, need_weights=False); small = small + c[12].unsqueeze(1) * u norm = block.norm_mlp(small) small = small + t[15].unsqueeze(1) * block.mlp(modulate(norm, t[13], t[14])) small = small + c[15].unsqueeze(1) * block.mlp(modulate(norm, c[13], c[14])) return self._decode(small, x) class AdaLNResidualLoRASWaG(AdaLNResidualSWaG): """Independent residual AdaLN condition chain plus cross-attention LoRA.""" def __init__(self, cross_attention_lora_rank=8, cross_attention_lora_alpha=16.0, cross_attention_lora_dropout=0.05, **kwargs): super().__init__(**kwargs) for block in self.blocks: for name in ("mid_attention", "large_attention"): base = getattr(block, name) adapted = LoRAMultiheadAttention(base.embed_dim, base.num_heads, cross_attention_lora_rank, cross_attention_lora_alpha, cross_attention_lora_dropout) adapted.in_proj_weight.data.copy_(base.in_proj_weight.data) adapted.in_proj_bias.data.copy_(base.in_proj_bias.data) adapted.out_proj.load_state_dict(base.out_proj.state_dict()) setattr(block, name, adapted)