Download training/models.py from DancingNow/swag-train-bundle: direct link, hf CLI and curl.
- Browser
- Download file 35 kB
-
https://huggingface.co/DancingNow/swag-train-bundle/resolve/main/training/models.py
- Command line
-
hf download hf://DancingNow/swag-train-bundle/training/models.py
-
curl -L -o models.py https://huggingface.co/DancingNow/swag-train-bundle/resolve/main/training/models.py
35 kB
| """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), | |
| ) | |
| 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) | |