DancingNow's picture
Add files using upload-large-folder tool
275b5a1 verified
Raw History Blame Contribute Delete
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),
)
@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)