import torch import torch.nn as nn import torch.nn.functional as F import numpy as np from einops import rearrange, repeat from torch.utils.checkpoint import checkpoint import math from typing import Optional, List, Callable from scipy.optimize import linear_sum_assignment # class FM: # def __init__(self, sigma_min=1e-5, timescale=1.0): # self.sigma_min = sigma_min # self.prediction_type = None # self.timescale = timescale # def alpha(self, t): # return 1.0 - t # def sigma(self, t): # return self.sigma_min + t * (1.0 - self.sigma_min) # def A(self, t): # return 1.0 # def B(self, t): # return -(1.0 - self.sigma_min) # def get_betas(self, n_timesteps): # return torch.zeros(n_timesteps) # Not VP and not supported # def add_noise(self, x, t, noise=None): # noise = torch.randn_like(x) if noise is None else noise # s = [x.shape[0], x.shape[1], x.shape[2], 1] # x_t = self.alpha(t).view(*s) * x + self.sigma(t).view(*s) * noise # return x_t, noise # def loss(self, net, x, t=None, net_kwargs=None, return_loss_unreduced=False, return_all=False): # B, T, N, C = x.shape # if net_kwargs is None: # net_kwargs = {} # if t is None: # t = torch.rand(B, T, device=x.device) # # t = torch.sigmoid(torch.randn(B, T, device=x.device)) # for logit normal # repeat_t = t.unsqueeze(2).repeat(1, 1, N) # x_t, noise = self.add_noise(x, repeat_t) # pred = net(x_t, t=t * self.timescale, **net_kwargs) # target = self.A(repeat_t) * x + self.B(repeat_t) * noise # -dxt/dt # if return_loss_unreduced: # loss = ((pred.float() - target.float()) ** 2).mean(dim=[1, 2]) # if return_all: # return loss, t, x_t, pred # else: # return loss, t # else: # loss = ((pred.float() - target.float()) ** 2).mean() # if return_all: # return loss, x_t, pred # else: # return loss # def get_prediction( # self, # net, # x_t, # t, # net_kwargs=None, # uncond_net_kwargs=None, # guidance=1.0, # ): # if net_kwargs is None: # net_kwargs = {} # if guidance != 1.0: # assert uncond_net_kwargs is not None # uncond_pred = net(x_t, t=t * self.timescale, **uncond_net_kwargs) # cond_pred = net(x_t, t=t * self.timescale, **net_kwargs) # pred = uncond_pred + guidance * (cond_pred - uncond_pred) # # if guidance != 1.0: # # assert uncond_net_kwargs is not None # # x_t = torch.cat([x_t, x_t], dim=0) # # t = torch.cat([t, t], dim=0) # # combined_kwargs = {} # # # we assume the keys match # # for k, v in net_kwargs.items(): # # combined_kwargs[k] = torch.cat([v, uncond_net_kwargs[k]], dim=0) # # combined_pred = net(x_t, t=t * self.timescale, **combined_kwargs) # # pred, uncond_pred = combined_pred.chunk(2, dim=0) # # pred = uncond_pred + guidance * (pred - uncond_pred) # else: # pred = net(x_t, t=t * self.timescale, **net_kwargs) # return pred # def convert_sample_prediction(self, x_t, t, pred): # M = torch.tensor([ # [self.alpha(t), self.sigma(t)], # [self.A(t), self.B(t)], # ], dtype=torch.float64) # M_inv = torch.linalg.inv(M) # sample_pred = M_inv[0, 0].item() * x_t + M_inv[0, 1].item() * pred # return sample_pred # class FMEulerSampler: # def __init__(self, diffusion): # self.diffusion = diffusion # def sample( # self, # net, # shape, # n_steps, # net_kwargs=None, # uncond_net_kwargs=None, # guidance=1.0, # noise=None, # ): # """ # Implements simple uniform noise sampling for bidirectional generation # """ # device = next(net.parameters()).device # x_t = torch.randn(shape, device=device) if noise is None else noise # t_steps = torch.linspace(1, 0, n_steps + 1, device=device) # with torch.no_grad(): # for i in range(n_steps): # t = t_steps[i].repeat(x_t.shape[0], x_t.shape[1]) # neg_v = self.diffusion.get_prediction( # net, # x_t, # t, # net_kwargs=net_kwargs, # uncond_net_kwargs=uncond_net_kwargs, # guidance=guidance, # ) # x_t = x_t + neg_v * (t_steps[i] - t_steps[i + 1]) # return x_t class FM: def __init__(self, sigma_min=1e-5, timescale=1.0): self.sigma_min = sigma_min self.prediction_type = None self.timescale = timescale def alpha(self, t): return 1.0 - t def sigma(self, t): return self.sigma_min + t * (1.0 - self.sigma_min) def A(self, t): return 1.0 def B(self, t): return -(1.0 - self.sigma_min) def get_betas(self, n_timesteps): return torch.zeros(n_timesteps) # Not VP and not supported def add_noise(self, x, t, noise=None): noise = torch.randn_like(x) if noise is None else noise s = [x.shape[0], x.shape[1], x.shape[2], 1] x_t = self.alpha(t).view(*s) * x + self.sigma(t).view(*s) * noise return x_t, noise @torch.compiler.disable def get_ot_noise(self, x: torch.Tensor, noise: torch.Tensor): B = x.shape[0] x_flat = x.view(B, -1).detach() noise_flat = noise.view(B, -1).detach() # cost_matrix = torch.cdist(x_flat, noise_flat, p=2) cost_matrix = torch.cdist(noise_flat, x_flat, p=2) cost_matrix_np = cost_matrix.cpu().numpy() _, col_ind = linear_sum_assignment(cost_matrix_np) col_ind_tensor = torch.from_numpy(col_ind).to(device=x.device, dtype=torch.long) return noise[col_ind_tensor] def loss(self, net, x, t=None, net_kwargs=None, return_loss_unreduced=False, return_all=False): B, T, N, C = x.shape if net_kwargs is None: net_kwargs = {} if t is None: # uniform # t = torch.rand(B, T, device=x.device) # logit normal t = torch.sigmoid(torch.randn(B, T, device=x.device)) repeat_t = t.unsqueeze(2).repeat(1, 1, N) noise = torch.randn_like(x) noise = self.get_ot_noise(x, noise) x_t, noise = self.add_noise(x, repeat_t, noise=noise) pred = net(x_t, t=t * self.timescale, **net_kwargs) target = self.A(repeat_t) * x + self.B(repeat_t) * noise # -dxt/dt if return_loss_unreduced: loss = ((pred.float() - target.float()) ** 2).mean(dim=[1, 2]) if return_all: return loss, t, x_t, pred else: return loss, t else: loss = ((pred.float() - target.float()) ** 2).mean() if return_all: return loss, x_t, pred else: return loss @staticmethod def _concat_kwargs(kwarg_list, dim=0): """ Recursively concatenates tensors in a list of identical-structure dictionaries. Safely handles nested dictionaries like 'unconditional_mask'. """ combined = {} for k in kwarg_list[0].keys(): val = kwarg_list[0][k] if isinstance(val, dict): # Recurse for nested dicts (e.g., the unconditional_mask) combined[k] = FM._concat_kwargs([kw[k] for kw in kwarg_list], dim=dim) elif isinstance(val, torch.Tensor): # Concatenate tensors along the batch dimension combined[k] = torch.cat([kw[k] for kw in kwarg_list], dim=dim) else: # For non-tensors (e.g., bool flags or strings), assume they are # constant across the batch and just copy the first one. combined[k] = val return combined def get_prediction( self, net, x_t, t, net_kwargs=None, uncond_net_kwargs=None, guidance=1.0, memory_efficient=False, rescale_phi=0.0, cfg_mode="independent", ): if net_kwargs is None: net_kwargs = {} # Normalize inputs to lists if isinstance(net_kwargs, dict): net_kwargs_list = [net_kwargs] guidance_list = [guidance] else: net_kwargs_list = net_kwargs guidance_list = guidance if isinstance(guidance, list) else [guidance] * len(net_kwargs_list) is_cfg = any(g != 1.0 for g in guidance_list) or len(net_kwargs_list) > 1 if not is_cfg: # Standard single pass (no CFG) return net(x_t, t=t * self.timescale, **net_kwargs_list[0]) assert uncond_net_kwargs is not None, "uncond_net_kwargs must be provided when using guidance." # ========================================== # MODE 1: JOINT CFG (Single combined pass) # ========================================== if cfg_mode == "joint": # 1. Merge all isolated conditions into one master kwargs dictionary joint_kwargs = {} for kw in net_kwargs_list: joint_kwargs.update(kw) # 2. Pick a single guidance scale (defaults to the first one in the list) g = guidance_list[0] if not memory_efficient: batched_x_t = torch.cat([x_t] * 2, dim=0) batched_t = torch.cat([t] * 2, dim=0) combined_kwargs = self._concat_kwargs([uncond_net_kwargs, joint_kwargs], dim=0) combined_pred = net(batched_x_t, t=batched_t * self.timescale, **combined_kwargs) uncond_pred, joint_pred = combined_pred.chunk(2, dim=0) else: uncond_pred = net(x_t, t=t * self.timescale, **uncond_net_kwargs) joint_pred = net(x_t, t=t * self.timescale, **joint_kwargs) # Standard CFG Formula pred = uncond_pred + g * (joint_pred - uncond_pred) reference_pred = joint_pred # The reference is just the unscaled joint prediction # ========================================== # MODE 2: INDEPENDENT CFG (Compositional) # ========================================== elif cfg_mode == "independent": if not memory_efficient: n_passes = 1 + len(net_kwargs_list) batched_x_t = torch.cat([x_t] * n_passes, dim=0) batched_t = torch.cat([t] * n_passes, dim=0) list_to_cat = [uncond_net_kwargs] + net_kwargs_list combined_kwargs = self._concat_kwargs(list_to_cat, dim=0) combined_pred = net(batched_x_t, t=batched_t * self.timescale, **combined_kwargs) preds = combined_pred.chunk(n_passes, dim=0) uncond_pred = preds[0] cond_preds = preds[1:] pred = uncond_pred.clone() reference_pred = uncond_pred.clone() for g, cp in zip(guidance_list, cond_preds): delta = cp - uncond_pred pred += g * delta reference_pred += delta else: uncond_pred = net(x_t, t=t * self.timescale, **uncond_net_kwargs) pred = uncond_pred.clone() reference_pred = uncond_pred.clone() for kwargs, g in zip(net_kwargs_list, guidance_list): if g == 0.0: continue cp = net(x_t, t=t * self.timescale, **kwargs) delta = cp - uncond_pred pred += g * delta reference_pred += delta else: raise ValueError(f"Unknown cfg_mode: {cfg_mode}") # ========================================== # GUIDANCE RESCALING (Applies to both modes) # ========================================== if rescale_phi > 0.0: dims_to_reduce = tuple(range(1, pred.ndim)) std_cfg = pred.std(dim=dims_to_reduce, keepdim=True) std_ref = reference_pred.std(dim=dims_to_reduce, keepdim=True) factor = std_ref / (std_cfg + 1e-8) pred_rescaled = pred * factor pred = rescale_phi * pred_rescaled + (1.0 - rescale_phi) * pred return pred def convert_sample_prediction(self, x_t, t, pred): M = torch.tensor([ [self.alpha(t), self.sigma(t)], [self.A(t), self.B(t)], ], dtype=torch.float64) M_inv = torch.linalg.inv(M) sample_pred = M_inv[0, 0].item() * x_t + M_inv[0, 1].item() * pred return sample_pred class FMEulerSampler: def __init__(self, diffusion): self.diffusion = diffusion def sample( self, net, shape, n_steps, net_kwargs=None, uncond_net_kwargs=None, guidance=1.0, noise=None, memory_efficient=False, rescale_phi=0, cfg_mode="independent", t_dist="uniform", ): """ Implements simple uniform noise sampling for bidirectional generation Supports Compositional CFG by passing lists to net_kwargs and guidance. """ assert t_dist in ['uniform', 'logit'], f't_dist must be uniform or logit but got {t_dist}' device = next(net.parameters()).device x_t = torch.randn(shape, device=device) if noise is None else noise if t_dist == 'uniform': t_steps = torch.linspace(1, 0, n_steps + 1, device=device) elif t_dist == 'logit': u = torch.linspace(1.0 - 1e-5, 1e-5, n_steps + 1, device=device) z = math.sqrt(2.0) * torch.erfinv(2.0 * u - 1.0) t_steps = torch.sigmoid(z) t_steps[0] = 1.0 t_steps[-1] = 0.0 with torch.no_grad(): for i in range(n_steps): t = t_steps[i].repeat(x_t.shape[0], x_t.shape[1]) neg_v = self.diffusion.get_prediction( net, x_t, t, net_kwargs=net_kwargs, uncond_net_kwargs=uncond_net_kwargs, guidance=guidance, memory_efficient=memory_efficient, rescale_phi=rescale_phi, cfg_mode=cfg_mode ) x_t = x_t + neg_v * (t_steps[i] - t_steps[i + 1]) return x_t # @torch.compile def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor): return x * (1 + scale) + shift def apply_scaling(freqs: torch.Tensor): # RoPE scaling (values obtained from grid search) scale_factor = 8 low_freq_factor = 1 high_freq_factor = 4 old_context_len = 8192 # original llama3 length low_freq_wavelen = old_context_len / low_freq_factor high_freq_wavelen = old_context_len / high_freq_factor new_freqs = [] for freq in freqs: wavelen = 2 * math.pi / freq if wavelen < high_freq_wavelen: new_freqs.append(freq) elif wavelen > low_freq_wavelen: new_freqs.append(freq / scale_factor) else: assert low_freq_wavelen != high_freq_wavelen smooth = (old_context_len / wavelen - low_freq_factor) / ( high_freq_factor - low_freq_factor ) new_freqs.append((1 - smooth) * freq / scale_factor + smooth * freq) return torch.tensor(new_freqs, dtype=freqs.dtype, device=freqs.device) def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0, use_scaled: bool = False): freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim)) t = torch.arange(end, device=freqs.device, dtype=torch.float32) if use_scaled: freqs = apply_scaling(freqs) freqs = torch.outer(t, freqs) freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # complex64 freqs_cis_real = torch.stack([freqs_cis.real, freqs_cis.imag], dim=-1) return freqs_cis_real def apply_rotary_emb(x, freqs_cis): # shape gymnastics let's go # x is (bs, seqlen, n_heads, head_dim), e.g. (4, 8, 32, 128) # freqs_cis is (seq_len, head_dim/2, 2), e.g. (8, 64, 2) xshaped = x.float().reshape(*x.shape[:-1], -1, 2) # xshaped is (bs, seqlen, n_heads, head_dim/2, 2), e.g. (4, 8, 32, 64, 2) freqs_cis = freqs_cis.view(1, xshaped.size(1), 1, xshaped.size(3), 2) # freqs_cis becomes (1, seqlen, 1, head_dim/2, 2), e.g. (1, 8, 1, 64, 2) x_out2 = torch.stack( [ xshaped[..., 0] * freqs_cis[..., 0] - xshaped[..., 1] * freqs_cis[..., 1], xshaped[..., 1] * freqs_cis[..., 0] + xshaped[..., 0] * freqs_cis[..., 1], ], -1, ) # x_out2 at this point is (bs, seqlen, n_heads, head_dim/2, 2), e.g. (4, 8, 32, 64, 2) x_out2 = x_out2.flatten(3) # x_out2 is now (bs, seqlen, n_heads, head_dim), e.g. (4, 8, 32, 128) return x_out2.type_as(x) class DropPath(nn.Module): """Stochastic Depth: Randomly drops paths (blocks) per sample during training.""" def __init__(self, drop_prob=0.0): super().__init__() self.drop_prob = drop_prob def forward(self, x): if self.drop_prob == 0. or not self.training: return x keep_prob = 1 - self.drop_prob shape = (x.shape[0],) + (1,) * (x.ndim - 1) random_tensor = keep_prob + torch.rand(shape, dtype=x.dtype, device=x.device) random_tensor.floor_() return x.div(keep_prob) * random_tensor class TimestepEmbedder(nn.Module): """ Embeds scalar timesteps into vector representations. """ def __init__(self, hidden_size, frequency_embedding_size=256, max_period=10000, bias=True, swiglu=False): super().__init__() if swiglu: self.mlp = SwiGLUMlp(frequency_embedding_size, int(2 / 3 * 4 * hidden_size), hidden_size, bias=bias) else: self.mlp = nn.Sequential( nn.Linear(frequency_embedding_size, hidden_size, bias=bias), nn.SiLU(), nn.Linear(hidden_size, hidden_size, bias=bias), ) self.frequency_embedding_size = frequency_embedding_size self.max_period = max_period @staticmethod def timestep_embedding(t, dim, max_period=10000): """ Create sinusoidal timestep embeddings. :param t: a 1-D Tensor of (N) indices, one per batch element. These may be fractional. :param dim: the dimension of the output. :param max_period: controls the minimum frequency of the embeddings. :return: an (N, D) Tensor of positional embeddings. """ # https://github.com/openai/glide-text2im/blob/main/glide_text2im/nn.py half = dim // 2 freqs = torch.exp( -math.log(max_period) * torch.arange(start=0, end=half, dtype=torch.float32) / half ).to(device=t.device) args = t[:, None].float() * freqs[None] embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1) if dim % 2: embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1) return embedding def forward(self, t): t_freq = self.timestep_embedding(t, self.frequency_embedding_size, max_period=self.max_period) t_emb = self.mlp(t_freq) return t_emb class RMSNorm(nn.Module): def __init__(self, dim, eps=1e-6): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) def _norm(self, x): return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) def forward(self, x): output = self._norm(x.float()).type_as(x) return output * self.weight class Attention(nn.Module): def __init__( self, dim: int, num_heads: int = 8, qkv_bias: bool = False, proj_bias: bool = True, attn_drop: float = 0., proj_drop: float = 0., ) -> None: super().__init__() assert dim % num_heads == 0, 'dim should be divisible by num_heads' self.num_heads = num_heads self.head_dim = dim // num_heads self.scale = self.head_dim ** -0.5 self.fused_attn = True self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias) self.attn_drop = nn.Dropout(attn_drop) self.proj = nn.Linear(dim, dim, bias=proj_bias) self.proj_drop = nn.Dropout(proj_drop) def forward( self, x: torch.Tensor, freqs_cis: Optional[torch.Tensor] = None, attn_mask = None, is_causal: bool = False ) -> torch.Tensor: B, N, C = x.shape qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4) q, k, v = qkv.unbind(0) # RoPE if freqs_cis is not None: q = apply_rotary_emb(q.transpose(1, 2), freqs_cis).transpose(1, 2) k = apply_rotary_emb(k.transpose(1, 2), freqs_cis).transpose(1, 2) if self.fused_attn: x = F.scaled_dot_product_attention( q, k, v, attn_mask=attn_mask, is_causal=is_causal, dropout_p=self.attn_drop.p if self.training else 0., ) else: raise NotImplementedError() q = q * self.scale attn = q @ k.transpose(-2, -1) attn = maybe_add_mask(attn, attn_mask) attn = attn.softmax(dim=-1) attn = self.attn_drop(attn) x = attn @ v x = x.transpose(1, 2).reshape(B, N, C) x = self.proj(x) x = self.proj_drop(x) return x class CrossAttention(nn.Module): """ Multi-head cross-attention module. Query: x Key/Value: context """ def __init__( self, dim: int, num_heads: int = 8, qkv_bias: bool = False, proj_bias: bool = True, attn_drop: float = 0., proj_drop: float = 0., ) -> None: super().__init__() assert dim % num_heads == 0, 'dim should be divisible by num_heads' self.num_heads = num_heads self.head_dim = dim // num_heads self.scale = self.head_dim ** -0.5 self.fused_attn = True # Separate linear layers for query (from x) and key/value (from context) self.q = nn.Linear(dim, dim, bias=qkv_bias) self.kv = nn.Linear(dim, dim * 2, bias=qkv_bias) self.attn_drop = nn.Dropout(attn_drop) self.proj = nn.Linear(dim, dim, bias=proj_bias) self.proj_drop = nn.Dropout(proj_drop) def forward( self, x: torch.Tensor, context: torch.Tensor, freqs_cis: Optional[torch.Tensor] = None, attn_mask = None, ) -> torch.Tensor: """ x: [B, N, C] query sequence context: [B, M, C] key/value sequence attn_mask: optional [B, N, M] mask """ B, N, C = x.shape _, M, _ = context.shape # Linear projections q = self.q(x).reshape(B, N, self.num_heads, self.head_dim).permute(0, 2, 1, 3) # [B, H, N, D] kv = self.kv(context).reshape(B, M, 2, self.num_heads, self.head_dim).permute(2, 0, 3, 1, 4) k, v = kv.unbind(0) # [B, H, M, D] # RoPE if freqs_cis is not None: q = apply_rotary_emb(q.transpose(1, 2), freqs_cis[:q.shape[2]]).transpose(1, 2) k = apply_rotary_emb(k.transpose(1, 2), freqs_cis[:k.shape[2]]).transpose(1, 2) if self.fused_attn: # PyTorch 2.1+ scaled_dot_product_attention supports cross-attention x = F.scaled_dot_product_attention( q, k, v, attn_mask=attn_mask, dropout_p=self.attn_drop.p if self.training else 0., ) else: raise NotImplementedError() # fallback: manual attention q = q * self.scale attn = q @ k.transpose(-2, -1) # [B, H, N, M] if attn_mask is not None: attn = attn.masked_fill(attn_mask.bool(), float('-inf')) attn = attn.softmax(dim=-1) attn = self.attn_drop(attn) x = attn @ v # [B, H, N, D] x = x.transpose(1, 2).reshape(B, N, C) x = self.proj(x) x = self.proj_drop(x) return x class SwiGLUMlp(nn.Module): def __init__( self, in_features: int, hidden_features: Optional[int] = None, out_features: Optional[int] = None, act_layer: Callable[..., nn.Module] = None, drop: float = 0.0, bias: bool = True, ) -> None: super().__init__() out_features = out_features or in_features hidden_features = hidden_features or in_features self.w12 = nn.Linear(in_features, 2 * hidden_features, bias=bias) self.w3 = nn.Linear(hidden_features, out_features, bias=bias) # @torch.compile def forward(self, x: torch.Tensor) -> torch.Tensor: x12 = self.w12(x) x1, x2 = x12.chunk(2, dim=-1) hidden = F.silu(x1) * x2 return self.w3(hidden) class DiTBlock(nn.Module): def __init__(self, hidden_size, num_heads, mlp_ratio=4.0, drop_path=0, **block_kwargs): super().__init__() self.norm1 = nn.LayerNorm(hidden_size) self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=False, proj_bias=True, **block_kwargs) self.norm2 = nn.LayerNorm(hidden_size) self.mlp = SwiGLUMlp(hidden_size, int(2 / 3 * mlp_ratio * hidden_size), bias=True) self.scale_shift_table = nn.Parameter( torch.randn(6, hidden_size) / hidden_size ** 0.5, ) self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity() def forward(self, x, t, freqs_cis=None, attn_mask=None): """ Incredibly ugly but trades huge memory savings for time """ B, TN, C = x.shape B, T, NC = t.shape N = TN // T biases = self.scale_shift_table[None, None] + t.reshape(x.size(0), T, 6, -1) ( shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp, ) = [chunk.expand(-1, -1, N, -1) for chunk in biases.chunk(6, dim=-2)] # ugly but memory saving... x = rearrange(x, 'b (t n) c -> b t n c', t=T, n=N) x = x + self.drop_path(gate_msa * rearrange(self.attn(rearrange(modulate(self.norm1(x), shift_msa, scale_msa), 'b t n c -> b (t n) c'), freqs_cis=freqs_cis, attn_mask=attn_mask), 'b (t n) c -> b t n c', t=T, n=N)) x = x + self.drop_path(gate_mlp * self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp))) x = rearrange(x, 'b t n c -> b (t n) c') return x class DiTAirBlock(nn.Module): def __init__(self, hidden_size, num_heads, mlp_ratio=4.0, drop_path=0, **block_kwargs): super().__init__() self.norm1 = nn.LayerNorm(hidden_size) self.attn = Attention(hidden_size, num_heads=num_heads, qkv_bias=False, proj_bias=True, **block_kwargs) self.norm2 = nn.LayerNorm(hidden_size) self.mlp = SwiGLUMlp(hidden_size, int(2 / 3 * mlp_ratio * hidden_size), bias=True) self.scale_shift_table = nn.Parameter( torch.randn(6, hidden_size) / hidden_size ** 0.5, ) self.drop_path = DropPath(drop_path) if drop_path > 0. else nn.Identity() def forward(self, x, t, freqs_cis=None, attn_mask=None): B, N, C = x.shape B, T, C = t.shape biases = self.scale_shift_table[None] + t.reshape(x.size(0), T, 6, -1) ( shift_msa_T, scale_msa_T, gate_msa_T, shift_mlp_T, scale_mlp_T, gate_mlp_T, ) = [chunk.squeeze(2) for chunk in biases.chunk(6, dim=-2)] shift_msa = torch.zeros_like(x) shift_msa[:, -T:] = shift_msa_T shift_mlp = torch.zeros_like(x) shift_mlp[:, -T:] = shift_mlp_T gate_msa = torch.zeros_like(x) gate_msa[:, -T:] = gate_msa_T gate_mlp = torch.zeros_like(x) gate_mlp[:, -T:] = gate_mlp_T scale_msa = torch.ones_like(x) scale_msa[:, -T:] = scale_msa_T scale_mlp = torch.ones_like(x) scale_mlp[:, -T:] = scale_mlp_T x = x + self.drop_path(gate_msa * self.attn(modulate(self.norm1(x), shift_msa, scale_msa), freqs_cis=freqs_cis, attn_mask=attn_mask)) x = x + self.drop_path(gate_mlp * self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp))) return x def create_block_causal_mask(block_size: int, num_blocks: int, dtype=torch.float32): """ Creates a block causal mask where tokens can attend to their own block and all previous blocks, but not future blocks. Args: block_size (int): The length of each block. num_blocks (int): The number of blocks. dtype: The data type for the mask (default: torch.float32). Returns: torch.Tensor: A mask of shape (seq_len, seq_len) where 0.0 indicates 'attend' and -inf indicates 'mask'. (seq_len = block_size * num_blocks) """ # 1. Create a vector of block IDs: [0, 0, ..., 1, 1, ..., 2, 2, ...] block_ids = torch.arange(num_blocks).repeat_interleave(block_size) # 2. Broadcast to create a grid of block comparisons # Shape becomes (seq_len, 1) and (1, seq_len) for broadcasting row_ids = block_ids.unsqueeze(1) col_ids = block_ids.unsqueeze(0) # 3. Create boolean mask: True if row_block >= col_block (Past or Current Block) # This allows full bidirectional attention within the block mask_bool = row_ids >= col_ids return mask_bool def token_drop(labels, null_token, training, p_uncond=0.1, p_full=0.3, p_ind_low=0.1, p_ind_high=0.6): """ Partitions the batch into three mutually exclusive training modes: 1. Unconditional (Drop All) 2. Full Conditional (Keep All) 3. Partial Conditional (Drop Individual Tokens) Args: labels: (B, ...) Input tensor null_token: (1, C) Learnable null vector p_uncond: Probability of the Unconditional mode. p_full: Probability of the Full Conditional mode. p_ind_drop: Probability of dropping a token *given* we are in Partial mode. """ if not training: return labels B = labels.shape[0] device = labels.device batch_rand_shape = (B,) + (1,) * (labels.ndim - 2) batch_rand = torch.rand(batch_rand_shape, device=device) mask_drop_all = batch_rand < p_uncond mask_partial_mode = batch_rand >= (p_uncond + p_full) sample_specific_drop_rates = torch.rand(batch_rand_shape, device=device) * (p_ind_high - p_ind_low) + p_ind_low token_noise = torch.rand(labels.shape[:-1], device=device) mask_token_drop = token_noise < sample_specific_drop_rates final_mask = mask_drop_all.unsqueeze(-1) | (mask_partial_mode.unsqueeze(-1) & mask_token_drop.unsqueeze(-1)) null_token = null_token.to(labels.dtype) return torch.where(final_mask, null_token, labels) # def multi_token_drop( # signals: dict, # null_tokens: dict, # training: bool, # p_joint_uncond=0.1, # p_joint_full=0.2, # p_ind_uncond=0.1, # p_ind_low=0.1, # p_ind_high=0.5 # ): # """ # Applies hierarchical CFG dropping across multiple conditioning signals. # Hierarchy: # 1. Joint Uncond (p_joint_uncond): Drops ALL signals for the batch item. # 2. Joint Full (p_joint_full): Keeps ALL signals perfectly intact. # 3. Independent Mode: For remaining batch items, each signal independently decides: # a) Independent Uncond (p_ind_uncond): Drop this specific signal entirely. # b) Partial Token Drop: Drop individual tokens within this signal sequence. # Args: # signals: Dict of input tensors, e.g., {'chroma': (B, T, C), 'bpm': (B, T, C)} # null_tokens: Dict of learnable null vectors, e.g., {'chroma': (1, C)} # training: Boolean flag # """ # if not training: # return signals # # Extract batch size and device from the first signal # first_key = list(signals.keys())[0] # B = signals[first_key].shape[0] # device = signals[first_key].device # # For a (B, T, C) tensor, this creates a shape of (B, 1) # batch_rand_shape = (B,) + (1,) * (signals[first_key].ndim - 2) # # --- LEVEL 1: SYNCHRONIZED JOINT MASKS --- # # These masks are shared across ALL signals to maintain the joint distribution # batch_rand = torch.rand(batch_rand_shape, device=device) # mask_joint_uncond = batch_rand < p_joint_uncond # mask_joint_full = batch_rand >= (1.0 - p_joint_full) # mask_independent_mode = ~(mask_joint_uncond | mask_joint_full) # output_signals = {} # # --- LEVEL 2: INDEPENDENT MASKS --- # for key, labels in signals.items(): # null_t = null_tokens[key].to(labels.dtype) # # 1. Independent Unconditional Drop (Drop the entire sequence for this specific signal) # ind_rand = torch.rand(batch_rand_shape, device=device) # mask_ind_uncond = mask_independent_mode & (ind_rand < p_ind_uncond) # # 2. Token-Level Drop (Only happens if we are in independent mode AND didn't drop the whole signal) # mask_token_mode = mask_independent_mode & ~mask_ind_uncond # # Randomize drop sparsity per batch item # sample_specific_drop_rates = torch.rand(batch_rand_shape, device=device) * (p_ind_high - p_ind_low) + p_ind_low # # Generate noise for every token in the sequence -> (B, T) # token_noise = torch.rand(labels.shape[:-1], device=device) # # Evaluate which tokens to drop # # sample_specific_drop_rates broadcasts from (B, 1) to (B, T) naturally # # mask_token_drop = token_noise < sample_specific_drop_rates.squeeze(-1) if sample_specific_drop_rates.dim() > 1 else token_noise < sample_specific_drop_rates # mask_token_drop = token_noise < sample_specific_drop_rates # # --- COMBINE ALL MASKS --- # # Final mask shape needs to be (B, T, 1) to broadcast over the channel dimension # final_mask = mask_joint_uncond.unsqueeze(-1) | \ # mask_ind_uncond.unsqueeze(-1) | \ # (mask_token_mode.unsqueeze(-1) & mask_token_drop.unsqueeze(-1)) # # Apply the mask: Replace dropped tokens with the specific null token for this signal # output_signals[key] = torch.where(final_mask, null_t, labels) # return output_signals def multi_token_drop( signals: dict, null_tokens: dict, training: bool, p_joint_uncond=0.10, p_joint_full=0.40, p_one_hot=0.30, p_ind_uncond=0.20, p_ind_low=0.05, p_ind_high=0.30, return_masks=False, ): """ Applies hierarchical CFG dropping optimized for Compositional CFG inference. Hierarchy (Batched Partitioning): 1. Joint Uncond (10%): Drops ALL signals. 2. Joint Full (40%): Keeps ALL signals pristine. 3. One-Hot Mode (30%): Keeps EXACTLY 1 signal, drops the rest. 4. Independent Mode (20%): Binomial dropping per signal, plus partial sequence drops. If return_masks=True, also returns a dict of per-signal boolean drop masks of shape (B, T) (True = the token was dropped/nulled). Useful for building presence channels for signals whose null value is not otherwise distinguishable. """ if not training: if return_masks: masks = { k: torch.zeros(v.shape[:2], dtype=torch.bool, device=v.device) for k, v in signals.items() } return signals, masks return signals first_key = list(signals.keys())[0] B = signals[first_key].shape[0] device = signals[first_key].device num_signals = len(signals) # --- LEVEL 0: BATCH PARTITIONING --- # Generate a single random float per batch item to route it to one of the 4 modes batch_rand = torch.rand((B,), device=device) limit_1 = p_joint_uncond limit_2 = limit_1 + p_joint_full limit_3 = limit_2 + p_one_hot mask_joint_uncond = batch_rand < limit_1 mask_joint_full = (batch_rand >= limit_1) & (batch_rand < limit_2) mask_one_hot = (batch_rand >= limit_2) & (batch_rand < limit_3) mask_independent = batch_rand >= limit_3 # Pre-calculate the "kept" index for the One-Hot slice of the batch # Each batch item in this mode will randomly select one index (0 to num_signals-1) to keep one_hot_keep_idx = torch.randint(0, num_signals, (B,), device=device) output_signals = {} drop_masks = {} # --- PROCESS EACH SIGNAL --- for idx, (key, labels) in enumerate(signals.items()): null_t = null_tokens[key].to(labels.dtype).to(device) T = labels.shape[1] # 1. One-Hot Drop Logic # Drop the signal if we are in one-hot mode AND it was not the randomly selected index is_not_the_kept_signal = (idx != one_hot_keep_idx) mask_drop_one_hot = mask_one_hot & is_not_the_kept_signal # 2. Independent Unconditional Drop Logic # Calculate a 20% drop chance, applied ONLY if routed to independent mode ind_rand = torch.rand((B,), device=device) mask_ind_uncond = mask_independent & (ind_rand < p_ind_uncond) # Combine all sequence-level dropping scenarios into one 1D mask: Shape (B,) full_drop_mask = mask_joint_uncond | mask_drop_one_hot | mask_ind_uncond # 3. Token-Level Drop Logic # Only applies to batch items in independent mode where the signal SURVIVED the un-cond drop mask_token_mode = mask_independent & ~mask_ind_uncond # Determine how severe the sequence masking is per batch item sample_drop_rates = torch.rand((B,), device=device) * (p_ind_high - p_ind_low) + p_ind_low # Generate token-level noise and evaluate: Shape (B, T) token_noise = torch.rand((B, T), device=device) # Unsqueeze sample_drop_rates to (B, 1) so it broadcasts smoothly across the T dimension mask_token_drop = token_noise < sample_drop_rates.unsqueeze(1) # Apply the token mode gate mask_token_drop = mask_token_mode.unsqueeze(1) & mask_token_drop # --- COMBINE MASKS & APPLY --- # Broadcast the 1D full drop mask to 2D, then combine with token drops: Shape (B, T) final_mask = full_drop_mask.unsqueeze(1) | mask_token_drop # Keep the (B, T) boolean drop mask before we align dims for torch.where drop_masks[key] = final_mask # Safely align dimensions for torch.where # Expands (B, T) into (B, T, 1) or (B, T, 1, 1) based on target label shape while final_mask.ndim < labels.ndim: final_mask = final_mask.unsqueeze(-1) # Swap the dropped tokens for the learned null vectors output_signals[key] = torch.where(final_mask, null_t, labels) if return_masks: return output_signals, drop_masks return output_signals class ConvBlock1d(nn.Module): def __init__( self, in_channels: int, out_channels: int, *, kernel_size: int = 3, stride: int = 1, dilation: int = 1, num_groups: int = 8, bias: bool = True, use_norm: bool = True, ) -> None: super().__init__() # Composer does not normalize raw conditioning before injecting it; when # use_norm=False we skip GroupNorm so channels (notably the binary presence # channels) aren't mixed together across the group. self.groupnorm = nn.GroupNorm( num_groups=num_groups, num_channels=in_channels ) if use_norm else nn.Identity() self.activation = nn.SiLU() self.project = nn.Conv1d( in_channels=in_channels, out_channels=out_channels, kernel_size=kernel_size, stride=stride, dilation=dilation, padding=kernel_size//2, bias=bias ) def forward( self, x: torch.Tensor, ) -> torch.Tensor: x = self.groupnorm(x) x = self.activation(x) return self.project(x) class ResnetBlock1d(nn.Module): def __init__( self, in_channels: int, out_channels: int, *, kernel_size: int = 3, stride: int = 1, dilation: int = 1, num_groups: int = 8, bias: bool = True, use_norm: bool = True, ) -> None: super().__init__() self.block1 = ConvBlock1d( in_channels=in_channels, out_channels=out_channels, kernel_size=kernel_size, stride=stride, dilation=dilation, num_groups=num_groups, bias=bias, use_norm=use_norm, ) self.block2 = ConvBlock1d( in_channels=out_channels, out_channels=out_channels, kernel_size=kernel_size, stride=1, dilation=dilation, num_groups=num_groups, bias=bias, use_norm=use_norm, ) if in_channels != out_channels or stride != 1: self.to_out = nn.Conv1d( in_channels=in_channels, out_channels=out_channels, kernel_size=1, stride=stride, bias=bias ) else: self.to_out = nn.Identity() def forward(self, x: torch.Tensor) -> torch.Tensor: h = self.block1(x) h = self.block2(h) return h + self.to_out(x) class Patcher(torch.nn.Module): def __init__( self, in_channels: int, out_channels: int, patch_size: int = 2, bias: bool = True, use_norm: bool = True, ): super().__init__() self.patch_size = patch_size self.block = ResnetBlock1d( in_channels=in_channels, out_channels=out_channels, stride=patch_size, num_groups=1, bias=bias, use_norm=use_norm, ) def forward(self, x: torch.Tensor) -> torch.Tensor: x = self.block(x) return x def zero_init_local_embedder(patcher: Patcher) -> None: """Make a Patcher emit exactly zero while staying trainable. ResnetBlock1d returns `block2(block1(x)) + to_out(x)`, so zeroing the two output-side projections is sufficient for a true no-op. block1 keeps its normal init on purpose -- see the note in ModernDiT.initialize_weights. """ for layer in (patcher.block.block2.project, patcher.block.to_out): nn.init.zeros_(layer.weight) if layer.bias is not None: nn.init.zeros_(layer.bias) class ModernDiT(nn.Module): def __init__(self, in_channels, hidden_size, spatial_window, n_chunks, style_dim, num_heads=12, depth=12, mlp_ratio=4, use_null_token=False, ): super().__init__() self.spatial_window = spatial_window self.use_null_token = use_null_token max_input_size = spatial_window * n_chunks self.t_embedder = TimestepEmbedder(hidden_size, bias=True, swiglu=True) self.bpm_embedder = TimestepEmbedder(hidden_size, bias=True, swiglu=True, max_period=1000) self.x_embedder = Patcher(in_channels, hidden_size) self.fuse_conditioning = SwiGLUMlp(hidden_size + style_dim, int(2 / 3 * mlp_ratio * hidden_size), hidden_size, bias=True) if self.use_null_token: self.null_token = nn.Parameter(torch.randn(style_dim) / style_dim ** 0.5) self.t_block = nn.Sequential( nn.SiLU(), nn.Linear(hidden_size, hidden_size * 6, bias=True), ) self.blocks = nn.ModuleList([ DiTBlock(hidden_size, num_heads, mlp_ratio=mlp_ratio) for _ in range(depth) ]) self.norm = nn.LayerNorm(hidden_size) self.final_layer_scale_shift_table = nn.Parameter( torch.randn(2, hidden_size) / hidden_size ** 0.5, ) self.fc = nn.Linear(hidden_size, in_channels, bias=False) self.initialize_weights() self.register_buffer('freqs_cis', precompute_freqs_cis(hidden_size // num_heads, max_input_size)) def initialize_weights(self): self.apply(self._init_weights) # zero out classifier weights nn.init.zeros_(self.fc.weight) nn.init.zeros_(self.t_block[-1].weight) nn.init.zeros_(self.t_block[-1].bias) # zero out c_proj weights in all blocks for block in self.blocks: nn.init.zeros_(block.mlp.w3.weight) nn.init.zeros_(block.attn.proj.weight) def _init_weights(self, module): if isinstance(module, nn.Linear): # https://arxiv.org/pdf/2310.17813 fan_out = module.weight.size(0) fan_in = module.weight.size(1) std = 1.0 / math.sqrt(fan_in) * min(1.0, math.sqrt(fan_out / fan_in)) nn.init.normal_(module.weight, mean=0.0, std=std) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=0.02) def forward(self, x, t, bpm, actions, unconditional_mask=None): B, T, N, C = x.shape x = rearrange(x, 'b t n c -> (b t) c n') x = self.x_embedder(x) x = rearrange(x, '(b t) c n -> b t n c', b=B, t=T) bpm = self.bpm_embedder(bpm.flatten()).view(B, T, 1, -1) x = x + bpm x = rearrange(x, 'b t n c -> b (t n) c') t = self.t_embedder(t.flatten()).view(B, T, -1) if self.use_null_token: actions = token_drop(actions, self.null_token.unsqueeze(0), self.training, p_uncond=0.1, p_full=0.8, p_ind_low=0.1, p_ind_high=0.5) if unconditional_mask is not None: actions = torch.where(unconditional_mask, self.null_token.unsqueeze(0), actions) t = torch.cat([t, actions], dim=-1) t = self.fuse_conditioning(t) t0 = self.t_block(t) freqs_cis = self.freqs_cis[:x.shape[1]] for block in self.blocks: x = block(x, t0, freqs_cis=freqs_cis) # SAM Audio does not use a non-linearity on t here shift, scale = (self.final_layer_scale_shift_table[None, None] + F.silu(t[:, :, None])).chunk( 2, dim=2 ) x = rearrange(x, 'b (t n) c -> b t n c', t=T, n=N) x = modulate(self.norm(x), shift.expand(-1, -1, N, -1), scale.expand(-1, -1, N, -1)) x = self.fc(x) return x class ModernDiTWrapper(nn.Module): def __init__(self, **kwargs): super().__init__() self.net = ModernDiT(**kwargs) self.diffusion = FM(timescale=1000.0) self.sampler = FMEulerSampler(self.diffusion) def forward(self, x, bpm, actions, t=None): return self.diffusion.loss(self.net, x, t=t, net_kwargs={'actions': actions, 'bpm': bpm}) def generate(self, x, bpm, actions, unconditional_mask=None, n_steps=50, uncond_net_kwargs=None, guidance=1.0): return self.sampler.sample(self.net, x.shape, n_steps=n_steps, net_kwargs={'actions': actions, 'bpm': bpm, 'unconditional_mask': unconditional_mask}, uncond_net_kwargs=uncond_net_kwargs, guidance=guidance) class UnconditionalModernDiT(nn.Module): def __init__(self, in_channels, hidden_size, spatial_window, n_chunks, num_heads=12, depth=12, mlp_ratio=4, gradient_checkpointing=False, patch_size=1, **kwargs, ): super().__init__() self.spatial_window = spatial_window self.gradient_checkpointing = gradient_checkpointing max_input_size = spatial_window * n_chunks self.patch_size = patch_size self.t_embedder = TimestepEmbedder(hidden_size, bias=False, swiglu=True) self.x_embedder = Patcher(in_channels, hidden_size, patch_size=patch_size) self.t_block = nn.Sequential( nn.SiLU(), nn.Linear(hidden_size, hidden_size * 6, bias=True), ) self.blocks = nn.ModuleList([ DiTBlock(hidden_size, num_heads, mlp_ratio=mlp_ratio) for _ in range(depth) ]) self.norm = RMSNorm(hidden_size) self.final_layer_scale_shift_table = nn.Parameter( torch.randn(2, hidden_size) / hidden_size ** 0.5, ) self.fc = nn.Linear(hidden_size, in_channels * patch_size, bias=False) self.initialize_weights() self.register_buffer('freqs_cis', precompute_freqs_cis(hidden_size // num_heads, max_input_size)) def initialize_weights(self): self.apply(self._init_weights) # zero out classifier weights nn.init.zeros_(self.fc.weight) nn.init.zeros_(self.t_block[-1].weight) nn.init.zeros_(self.t_block[-1].bias) # zero out c_proj weights in all blocks for block in self.blocks: nn.init.zeros_(block.mlp.w3.weight) nn.init.zeros_(block.attn.proj.weight) def _init_weights(self, module): if isinstance(module, nn.Linear): # https://arxiv.org/pdf/2310.17813 fan_out = module.weight.size(0) fan_in = module.weight.size(1) std = 1.0 / math.sqrt(fan_in) * min(1.0, math.sqrt(fan_out / fan_in)) nn.init.normal_(module.weight, mean=0.0, std=std) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=0.02) def forward(self, x, t): B, T, N, C = x.shape x = rearrange(x, 'b t n c -> (b t) c n') x = self.x_embedder(x) x = rearrange(x, '(b t) c n -> b (t n) c', b=B, t=T) t = self.t_embedder(t.flatten()).view(B, T, -1) t0 = self.t_block(t) freqs_cis = self.freqs_cis[:x.shape[1]] for block in self.blocks: if self.gradient_checkpointing and self.training: x = checkpoint(block, x, t0, freqs_cis=freqs_cis, use_reentrant=False) else: x = block(x, t0, freqs_cis=freqs_cis) # SAM Audio does not use a non-linearity on t here shift, scale = (self.final_layer_scale_shift_table[None, None] + F.silu(t[:, :, None])).chunk( 2, dim=2 ) x = rearrange(x, 'b (t n) c -> b t n c', t=T, n=N // self.patch_size) x = modulate(self.norm(x), shift.expand(-1, -1, N // self.patch_size, -1), scale.expand(-1, -1, N // self.patch_size, -1)) x = self.fc(x) x = rearrange(x, 'b t n (p c) -> b t (n p) c', p=self.patch_size, c=C) return x class UnconditionalModernDiTWrapper(nn.Module): def __init__(self, **kwargs): super().__init__() self.net = UnconditionalModernDiT(**kwargs) self.diffusion = FM(timescale=1000.0) self.sampler = FMEulerSampler(self.diffusion) def forward(self, x, t=None): return self.diffusion.loss(self.net, x, t=t) def generate(self, shape, net_kwargs=None, uncond_net_kwargs=None, n_steps=50, guidance=1.0, noise=None, memory_efficient=True, rescale_phi=0, cfg_mode="independent", t_dist="uniform"): return self.sampler.sample( self.net, shape, n_steps=n_steps, net_kwargs=net_kwargs, uncond_net_kwargs=uncond_net_kwargs, guidance=guidance, noise=noise, memory_efficient=memory_efficient, rescale_phi=rescale_phi, cfg_mode=cfg_mode, t_dist=t_dist ) class StyleConditionalModernDiT(nn.Module): def __init__(self, in_channels, hidden_size, spatial_window, n_chunks, style_dim, num_heads=12, depth=12, mlp_ratio=4, gradient_checkpointing=False, patch_size=1, use_null_token=False, **kwargs, ): super().__init__() self.spatial_window = spatial_window self.gradient_checkpointing = gradient_checkpointing max_input_size = spatial_window * n_chunks self.patch_size = patch_size self.use_null_token = use_null_token self.t_embedder = TimestepEmbedder(hidden_size, bias=False, swiglu=True) self.x_embedder = Patcher(in_channels, hidden_size, patch_size=patch_size) self.c_embedder = nn.Linear(style_dim, hidden_size, bias=True) self.fuse_conditioning = SwiGLUMlp(hidden_size + hidden_size, int(2 / 3 * mlp_ratio * hidden_size), hidden_size, bias=False) if self.use_null_token: self.null_token = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.t_block = nn.Sequential( nn.SiLU(), nn.Linear(hidden_size, hidden_size * 6, bias=True), ) self.blocks = nn.ModuleList([ DiTBlock(hidden_size, num_heads, mlp_ratio=mlp_ratio) for _ in range(depth) ]) self.norm = RMSNorm(hidden_size) self.final_layer_scale_shift_table = nn.Parameter( torch.randn(2, hidden_size) / hidden_size ** 0.5, ) self.fc = nn.Linear(hidden_size, in_channels * patch_size, bias=False) self.initialize_weights() self.register_buffer('freqs_cis', precompute_freqs_cis(hidden_size // num_heads, max_input_size)) def initialize_weights(self): self.apply(self._init_weights) # zero out classifier weights nn.init.zeros_(self.fc.weight) nn.init.zeros_(self.t_block[-1].weight) nn.init.zeros_(self.t_block[-1].bias) # zero out c_proj weights in all blocks for block in self.blocks: nn.init.zeros_(block.mlp.w3.weight) nn.init.zeros_(block.attn.proj.weight) def _init_weights(self, module): if isinstance(module, nn.Linear): # https://arxiv.org/pdf/2310.17813 fan_out = module.weight.size(0) fan_in = module.weight.size(1) std = 1.0 / math.sqrt(fan_in) * min(1.0, math.sqrt(fan_out / fan_in)) nn.init.normal_(module.weight, mean=0.0, std=std) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=0.02) def forward(self, x, t, c, unconditional_mask=None): B, T, N, C = x.shape x = rearrange(x, 'b t n c -> (b t) c n') x = self.x_embedder(x) x = rearrange(x, '(b t) c n -> b (t n) c', b=B, t=T) t = self.t_embedder(t.flatten()).view(B, T, -1) c = self.c_embedder(c) if self.use_null_token: c = token_drop(c, self.null_token.unsqueeze(0), self.training, p_uncond=0.1, p_full=0.8, p_ind_low=0.1, p_ind_high=0.5) if unconditional_mask is not None: c = torch.where(unconditional_mask, self.null_token.unsqueeze(0), c) t = torch.cat([t, c], dim=-1) t = self.fuse_conditioning(t) t0 = self.t_block(t) freqs_cis = self.freqs_cis[:x.shape[1]] for block in self.blocks: if self.gradient_checkpointing and self.training: x = checkpoint(block, x, t0, freqs_cis=freqs_cis, use_reentrant=False) else: x = block(x, t0, freqs_cis=freqs_cis) # SAM Audio does not use a non-linearity on t here shift, scale = (self.final_layer_scale_shift_table[None, None] + F.silu(t[:, :, None])).chunk( 2, dim=2 ) x = rearrange(x, 'b (t n) c -> b t n c', t=T, n=N // self.patch_size) x = modulate(self.norm(x), shift.expand(-1, -1, N // self.patch_size, -1), scale.expand(-1, -1, N // self.patch_size, -1)) x = self.fc(x) x = rearrange(x, 'b t n (p c) -> b t (n p) c', p=self.patch_size, c=C) return x class StyleConditionalModernDiTWrapper(nn.Module): def __init__(self, **kwargs): super().__init__() self.net = StyleConditionalModernDiT(**kwargs) self.diffusion = FM(timescale=1000.0) self.sampler = FMEulerSampler(self.diffusion) def forward(self, x, c, t=None): return self.diffusion.loss(self.net, x, t=t, net_kwargs={'c': c}) def generate(self, shape, net_kwargs=None, uncond_net_kwargs=None, n_steps=50, guidance=1.0, noise=None): return self.sampler.sample(self.net, shape, n_steps=n_steps, net_kwargs=net_kwargs, uncond_net_kwargs=uncond_net_kwargs, guidance=guidance, noise=noise) class BpmRmsChromaStyleConditionalModernDiT(nn.Module): def __init__(self, in_channels, hidden_size, spatial_window, n_chunks, style_dim, num_heads=12, depth=12, mlp_ratio=4, gradient_checkpointing=False, patch_size=1, use_null_token=False, **kwargs, ): super().__init__() self.spatial_window = spatial_window self.gradient_checkpointing = gradient_checkpointing max_input_size = spatial_window * n_chunks self.patch_size = patch_size self.use_null_token = use_null_token self.t_embedder = TimestepEmbedder(hidden_size, bias=False, swiglu=True) self.x_embedder = Patcher(in_channels, hidden_size, patch_size=patch_size) self.style_embedder = nn.Linear(style_dim, hidden_size, bias=True) self.chroma_embedder = nn.Linear(12, hidden_size, bias=True) self.rms_embedder = TimestepEmbedder(hidden_size, bias=False, swiglu=True, max_period=10) self.bpm_embedder = TimestepEmbedder(hidden_size, bias=False, swiglu=True, max_period=10) self.measure_embedder = nn.Embedding(n_chunks, hidden_size) self.fuse_conditioning = SwiGLUMlp(hidden_size, int(2 / 3 * mlp_ratio * hidden_size), hidden_size, bias=False) if self.use_null_token: self.null_style = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.null_chroma = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.null_rms = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.null_bpm = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.t_block = nn.Sequential( nn.SiLU(), nn.Linear(hidden_size, hidden_size * 6, bias=True), ) self.blocks = nn.ModuleList([ DiTBlock(hidden_size, num_heads, mlp_ratio=mlp_ratio) for _ in range(depth) ]) self.norm = RMSNorm(hidden_size) self.final_layer_scale_shift_table = nn.Parameter( torch.randn(2, hidden_size) / hidden_size ** 0.5, ) self.fc = nn.Linear(hidden_size, in_channels * patch_size, bias=False) self.initialize_weights() self.register_buffer('freqs_cis', precompute_freqs_cis(hidden_size // num_heads, max_input_size)) def initialize_weights(self): self.apply(self._init_weights) # zero out classifier weights nn.init.zeros_(self.fc.weight) nn.init.zeros_(self.t_block[-1].weight) nn.init.zeros_(self.t_block[-1].bias) # zero out c_proj weights in all blocks for block in self.blocks: nn.init.zeros_(block.mlp.w3.weight) nn.init.zeros_(block.attn.proj.weight) def _init_weights(self, module): if isinstance(module, nn.Linear): # https://arxiv.org/pdf/2310.17813 fan_out = module.weight.size(0) fan_in = module.weight.size(1) std = 1.0 / math.sqrt(fan_in) * min(1.0, math.sqrt(fan_out / fan_in)) nn.init.normal_(module.weight, mean=0.0, std=std) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=0.02) def forward(self, x, t, bpm, rms, chroma, style, unconditional_mask=None): B, T, N, C = x.shape x = rearrange(x, 'b t n c -> (b t) c n') x = self.x_embedder(x) x = rearrange(x, '(b t) c n -> b t n c', b=B, t=T) measure_ids = torch.arange(T, device=x.device) measure_embs = self.measure_embedder(measure_ids).unsqueeze(1) x = x + measure_embs x = rearrange(x, 'b t n c -> b (t n) c', b=B, t=T) rms = (rms - 0.09749302) / 0.047287412 bpm = (bpm - 187.5) / (226.4151001 - 144.57830811) # IQR t = self.t_embedder(t.flatten()).view(B, T, -1) style = self.style_embedder(style) chroma = self.chroma_embedder(chroma) rms = self.rms_embedder(rms.flatten()).view(B, T, -1) bpm = self.bpm_embedder(bpm.flatten()).view(B, T, -1) if self.use_null_token: signals = {'style': style, 'chroma': chroma, 'bpm': bpm, 'rms': rms} null_tokens = {'style': self.null_style, 'chroma': self.null_chroma, 'bpm': self.null_bpm, 'rms': self.null_rms} signals = multi_token_drop(signals, null_tokens, self.training, p_ind_uncond=0.1, p_joint_full=0.8, p_ind_low=0.1, p_ind_high=0.5) style = signals['style'] chroma = signals['chroma'] bpm = signals['bpm'] rms = signals['rms'] if unconditional_mask is not None: style = torch.where(unconditional_mask['style'], self.null_style, style) chroma = torch.where(unconditional_mask['chroma'], self.null_chroma, chroma) bpm = torch.where(unconditional_mask['bpm'], self.null_bpm, bpm) rms = torch.where(unconditional_mask['rms'], self.null_rms, rms) t = t + style + chroma + rms + bpm t = self.fuse_conditioning(t) t0 = self.t_block(t) freqs_cis = self.freqs_cis[:x.shape[1]] for block in self.blocks: if self.gradient_checkpointing and self.training: x = checkpoint(block, x, t0, freqs_cis=freqs_cis, use_reentrant=False) else: x = block(x, t0, freqs_cis=freqs_cis) # SAM Audio does not use a non-linearity on t here shift, scale = (self.final_layer_scale_shift_table[None, None] + F.silu(t[:, :, None])).chunk( 2, dim=2 ) x = rearrange(x, 'b (t n) c -> b t n c', t=T, n=N // self.patch_size) x = modulate(self.norm(x), shift.expand(-1, -1, N // self.patch_size, -1), scale.expand(-1, -1, N // self.patch_size, -1)) x = self.fc(x) x = rearrange(x, 'b t n (p c) -> b t (n p) c', p=self.patch_size, c=C) return x class BpmRmsChromaStyleConditionalModernDiTWrapper(nn.Module): def __init__(self, **kwargs): super().__init__() self.net = BpmRmsChromaStyleConditionalModernDiT(**kwargs) self.diffusion = FM(timescale=1000.0) self.sampler = FMEulerSampler(self.diffusion) def forward(self, x, bpm, rms, chroma, style, t=None): return self.diffusion.loss(self.net, x, t=t, net_kwargs={'style': style, 'chroma': chroma, 'bpm': bpm, 'rms': rms}) def generate(self, shape, net_kwargs=None, uncond_net_kwargs=None, n_steps=50, guidance=1.0, noise=None, memory_efficient=False): return self.sampler.sample(self.net, shape, n_steps=n_steps, net_kwargs=net_kwargs, uncond_net_kwargs=uncond_net_kwargs, guidance=guidance, noise=noise, memory_efficient=memory_efficient) class MetaConditionalModernDiT(nn.Module): def __init__(self, in_channels, hidden_size, spatial_window, n_chunks, style_dim, num_heads=12, depth=12, mlp_ratio=4, gradient_checkpointing=False, patch_size=1, use_null_token=False, **kwargs, ): super().__init__() self.spatial_window = spatial_window self.gradient_checkpointing = gradient_checkpointing max_input_size = spatial_window * n_chunks self.patch_size = patch_size self.use_null_token = use_null_token self.t_embedder = TimestepEmbedder(hidden_size, bias=False, swiglu=True) self.x_embedder = Patcher(in_channels, hidden_size, patch_size=patch_size) self.style_embedder = nn.Linear(style_dim, hidden_size, bias=True) self.chroma_embedder = nn.Linear(12, hidden_size, bias=True) self.rms_embedder = TimestepEmbedder(hidden_size, bias=False, swiglu=True)#, max_period=20) self.bpm_embedder = TimestepEmbedder(hidden_size, bias=False, swiglu=True)#, max_period=20) # self.rms_embedder = nn.Sequential(nn.Linear(1, hidden_size, bias=True), SwiGLUMlp(hidden_size, int(2 / 3 * mlp_ratio * hidden_size), bias=False)) # self.bpm_embedder = nn.Sequential(nn.Linear(1, hidden_size, bias=True), SwiGLUMlp(hidden_size, int(2 / 3 * mlp_ratio * hidden_size), bias=False)) self.mfcc_embedder = nn.Linear(12, hidden_size, bias=True) self.density_embedder = TimestepEmbedder(hidden_size, bias=False, swiglu=True)#, max_period=20) self.zcr_embedder = TimestepEmbedder(hidden_size, bias=False, swiglu=True)#, max_period=20) # self.density_embedder = nn.Sequential(nn.Linear(1, hidden_size, bias=True), SwiGLUMlp(hidden_size, int(2 / 3 * mlp_ratio * hidden_size), bias=False)) # self.zcr_embedder = nn.Sequential(nn.Linear(1, hidden_size, bias=True), SwiGLUMlp(hidden_size, int(2 / 3 * mlp_ratio * hidden_size), bias=False)) self.measure_embedder = nn.Embedding(n_chunks, hidden_size) self.fuse_conditioning = SwiGLUMlp(hidden_size * 2, int(2 / 3 * mlp_ratio * hidden_size * 2), hidden_size, bias=False) if self.use_null_token: self.null_style = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.null_chroma = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.null_rms_low = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.null_rms_mid = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.null_rms_high = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.null_bpm = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.null_mfcc = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.null_density = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.null_zcr = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.t_block = nn.Sequential( nn.SiLU(), nn.Linear(hidden_size, hidden_size * 6, bias=True), ) self.blocks = nn.ModuleList([ DiTBlock(hidden_size, num_heads, mlp_ratio=mlp_ratio) for _ in range(depth) ]) self.norm = RMSNorm(hidden_size) self.final_layer_scale_shift_table = nn.Parameter( torch.randn(2, hidden_size) / hidden_size ** 0.5, ) self.fc = nn.Linear(hidden_size, in_channels * patch_size, bias=False) self.initialize_weights() self.register_buffer('freqs_cis', precompute_freqs_cis(hidden_size // num_heads, max_input_size)) self.register_buffer('mcff_mean', torch.tensor([ 113.30053, -17.395779, 27.279049, -11.116686, 3.1354604, -9.138969, -2.866072, -7.1674404, -1.6265253, -5.047512, -1.7705443, -5.0958815 ])) self.register_buffer('mfcc_std', torch.tensor([ 38.435783, 28.687775, 18.932358, 14.646409, 13.498735, 10.035576, 9.510887, 8.25433, 8.212691, 7.155225, 7.3324447, 6.5340915 ])) def initialize_weights(self): self.apply(self._init_weights) # zero out classifier weights nn.init.zeros_(self.fc.weight) nn.init.zeros_(self.t_block[-1].weight) nn.init.zeros_(self.t_block[-1].bias) # zero out c_proj weights in all blocks for block in self.blocks: nn.init.zeros_(block.mlp.w3.weight) nn.init.zeros_(block.attn.proj.weight) def _init_weights(self, module): if isinstance(module, nn.Linear): # https://arxiv.org/pdf/2310.17813 fan_out = module.weight.size(0) fan_in = module.weight.size(1) std = 1.0 / math.sqrt(fan_in) * min(1.0, math.sqrt(fan_out / fan_in)) nn.init.normal_(module.weight, mean=0.0, std=std) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=0.02) def forward(self, x, t, bpm, rms_low, rms_mid, rms_high, density, zcr, mfcc, chroma, style, unconditional_mask=None): B, T, N, C = x.shape x = rearrange(x, 'b t n c -> (b t) c n') x = self.x_embedder(x) x = rearrange(x, '(b t) c n -> b t n c', b=B, t=T) measure_ids = torch.arange(T, device=x.device) measure_embs = self.measure_embedder(measure_ids).unsqueeze(1) x = x + measure_embs x = rearrange(x, 'b t n c -> b (t n) c', b=B, t=T) # rms_low = (rms_low - 6.2784767) / 4.1725345 # rms_mid = (rms_mid - 3.2565875) / 1.7880434 # rms_high = (rms_high - 0.26109472) / 0.3474748 # density = (density - 2.5229013) / 1.230155 # zcr = (zcr - 0.10766766) / 0.048143145 # bpm = (bpm - 187.5) / (226.4151001 - 144.57830811) # IQR mfcc = (mfcc - self.mcff_mean) / self.mfcc_std t = self.t_embedder(t.flatten()).view(B, T, -1) style = self.style_embedder(style) chroma = self.chroma_embedder(chroma) mfcc = self.mfcc_embedder(mfcc) rms_low = self.rms_embedder(rms_low.flatten()).view(B, T, -1) rms_mid = self.rms_embedder(rms_mid.flatten()).view(B, T, -1) rms_high = self.rms_embedder(rms_high.flatten()).view(B, T, -1) bpm = self.bpm_embedder(bpm.flatten()).view(B, T, -1) density = self.density_embedder(density.flatten()).view(B, T, -1) zcr = self.zcr_embedder(zcr.flatten()).view(B, T, -1) # rms_low = self.rms_embedder(rms_low.unsqueeze(-1)) # rms_mid = self.rms_embedder(rms_mid.unsqueeze(-1)) # rms_high = self.rms_embedder(rms_high.unsqueeze(-1)) # bpm = self.bpm_embedder(bpm.unsqueeze(-1)) # density = self.density_embedder(density.unsqueeze(-1)) # zcr = self.zcr_embedder(zcr.unsqueeze(-1)) if self.use_null_token: signals = { 'style': style, 'chroma': chroma, 'bpm': bpm, 'rms_low': rms_low, 'rms_mid': rms_mid, 'rms_high': rms_high, 'mfcc': mfcc, 'density': density, 'zcr': zcr } null_tokens = { 'style': self.null_style, 'chroma': self.null_chroma, 'bpm': self.null_bpm, 'rms_low': self.null_rms_low, 'rms_mid': self.null_rms_mid, 'rms_high': self.null_rms_high, 'mfcc': self.null_mfcc, 'density': self.null_density, 'zcr': self.null_zcr } signals = multi_token_drop( signals, null_tokens, self.training, p_joint_uncond=0.1, p_joint_full=0.5, p_one_hot=0.3, p_ind_uncond=0.1, p_ind_low=0.05, p_ind_high=0.3 ) style = signals['style'] chroma = signals['chroma'] bpm = signals['bpm'] rms_low = signals['rms_low'] rms_mid = signals['rms_mid'] rms_high = signals['rms_high'] mfcc = signals['mfcc'] density = signals['density'] zcr = signals['zcr'] if unconditional_mask is not None: style = torch.where(unconditional_mask['style'], self.null_style, style) chroma = torch.where(unconditional_mask['chroma'], self.null_chroma, chroma) bpm = torch.where(unconditional_mask['bpm'], self.null_bpm, bpm) rms_low = torch.where(unconditional_mask['rms_low'], self.null_rms_low, rms_low) rms_mid = torch.where(unconditional_mask['rms_mid'], self.null_rms_mid, rms_mid) rms_high = torch.where(unconditional_mask['rms_high'], self.null_rms_high, rms_high) mfcc = torch.where(unconditional_mask['mfcc'], self.null_mfcc, mfcc) density = torch.where(unconditional_mask['density'], self.null_density, density) zcr = torch.where(unconditional_mask['zcr'], self.null_zcr, zcr) c = style + chroma + rms_low + rms_mid + rms_high + mfcc + density + zcr + bpm t = torch.cat([t, c], dim=-1) t = self.fuse_conditioning(t) t0 = self.t_block(t) freqs_cis = self.freqs_cis[:x.shape[1]] for block in self.blocks: if self.gradient_checkpointing and self.training: x = checkpoint(block, x, t0, freqs_cis=freqs_cis, use_reentrant=False) else: x = block(x, t0, freqs_cis=freqs_cis) # SAM Audio does not use a non-linearity on t here shift, scale = (self.final_layer_scale_shift_table[None, None] + F.silu(t[:, :, None])).chunk( 2, dim=2 ) x = rearrange(x, 'b (t n) c -> b t n c', t=T, n=N // self.patch_size) x = modulate(self.norm(x), shift.expand(-1, -1, N // self.patch_size, -1), scale.expand(-1, -1, N // self.patch_size, -1)) x = self.fc(x) x = rearrange(x, 'b t n (p c) -> b t (n p) c', p=self.patch_size, c=C) return x class MetaConditionalModernDiTWrapper(nn.Module): def __init__(self, **kwargs): super().__init__() self.net = MetaConditionalModernDiT(**kwargs) self.diffusion = FM(timescale=1000.0) self.sampler = FMEulerSampler(self.diffusion) def forward(self, x, bpm, rms_low, rms_mid, rms_high, density, zcr, mfcc, chroma, style, t=None): return self.diffusion.loss( self.net, x, t=t, net_kwargs={ 'style': style, 'chroma': chroma, 'bpm': bpm, 'rms_low': rms_low, 'rms_mid': rms_mid, 'rms_high': rms_high, 'density': density, 'zcr': zcr, 'mfcc': mfcc, } ) def generate(self, shape, net_kwargs=None, uncond_net_kwargs=None, n_steps=50, guidance=1.0, noise=None, memory_efficient=True, rescale_phi=0, cfg_mode="independent"): return self.sampler.sample( self.net, shape, n_steps=n_steps, net_kwargs=net_kwargs, uncond_net_kwargs=uncond_net_kwargs, guidance=guidance, noise=noise, memory_efficient=memory_efficient, rescale_phi=rescale_phi, cfg_mode=cfg_mode ) class MetaConditionalModernDiTV2(nn.Module): def __init__(self, in_channels, hidden_size, spatial_window, n_chunks, style_dim, num_heads=12, depth=12, mlp_ratio=4, gradient_checkpointing=False, patch_size=1, use_null_token=False, stage=1, drop_path_rate=0.1, **kwargs, ): super().__init__() self.spatial_window = spatial_window self.gradient_checkpointing = gradient_checkpointing max_input_size = spatial_window * n_chunks self.patch_size = patch_size self.use_null_token = use_null_token self.t_embedder = TimestepEmbedder(hidden_size, bias=False, swiglu=True) self.x_embedder = Patcher(in_channels, hidden_size, patch_size=patch_size, bias=True) # 16 value channels (12 chroma + rms + density + zcr + flatness) plus 5 presence-mask channels (1 shared for chroma, 1 each for the 4 scalars). self.local_embedder = Patcher(16 + 5, hidden_size, patch_size=1, bias=True, use_norm=False) self.style_embedder = nn.Linear(style_dim, hidden_size, bias=True) self.bpm_embedder = nn.Embedding(350, hidden_size) if self.use_null_token: self.null_style = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.null_bpm = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.t_block = nn.Sequential( nn.SiLU(), nn.Linear(hidden_size, hidden_size * 6, bias=True), ) dp_rates=[x.item() for x in torch.linspace(0, drop_path_rate, depth)] self.blocks = nn.ModuleList([ DiTBlock(hidden_size, num_heads, mlp_ratio=mlp_ratio, drop_path=dp_rates[i]) for i in range(depth) ]) self.norm = RMSNorm(hidden_size) self.final_layer_scale_shift_table = nn.Parameter( torch.randn(2, hidden_size) / hidden_size ** 0.5, ) self.fc = nn.Linear(hidden_size, in_channels * patch_size, bias=False) self.bias = nn.Parameter(torch.zeros(in_channels * patch_size)) self.register_buffer('freqs_cis', precompute_freqs_cis(hidden_size // num_heads, max_input_size)) self.register_buffer('chroma_mean', torch.tensor([ 0.45533183, 0.39680213, 0.44615716, 0.42044115, 0.40855545, 0.45450154, 0.3971631, 0.496346, 0.44164586, 0.4416672, 0.44793198, 0.39493898 ])) self.register_buffer('chroma_std', torch.tensor([ 0.18241853, 0.16477719, 0.18014704, 0.18011539, 0.1677363, 0.18919244, 0.16196373, 0.19185093, 0.18003348, 0.1768027, 0.18706752, 0.1618064 ])) self.register_buffer('rms_mean', torch.tensor([3.2653894])) self.register_buffer('rms_std', torch.tensor([3.597796])) self.register_buffer('density_mean', torch.tensor([2.5229013])) self.register_buffer('density_std', torch.tensor([1.230155])) self.register_buffer('zcr_mean', torch.tensor([0.10766766])) self.register_buffer('zcr_std', torch.tensor([0.048143145])) self.register_buffer('flatness_mean', torch.tensor([0.011151944])) self.register_buffer('flatness_std', torch.tensor([0.018700112])) self.initialize_weights() self.set_training_stage(stage) def set_training_stage(self, stage): assert stage in [1, 2], f'Stage must be 1 or 2 but got {stage}' self.stage = stage if self.stage == 1: for param in self.local_embedder.parameters(): param.requires_grad = False elif self.stage == 2: for param in self.local_embedder.parameters(): param.requires_grad = True def create_optimizer_groups(self, weight_decay=1e-2, base_lr=1e-4, new_lr=1e-3): base_decay = [] base_no_decay = [] new_decay = [] new_no_decay = [] for name, param in self.named_parameters(): if not param.requires_grad: continue new_layer = 'local_embedder' in name no_decay = param.ndim < 2 or name == 'bpm_embedder.weight' if new_layer: if no_decay: new_no_decay.append(param) else: new_decay.append(param) else: if no_decay: base_no_decay.append(param) else: base_decay.append(param) optim_groups = [ {"params": base_decay, "weight_decay": weight_decay, "lr": base_lr}, {"params": base_no_decay, "weight_decay": 0.0, "lr": base_lr}, {"params": new_decay, "weight_decay": weight_decay, "lr": new_lr}, {"params": new_no_decay, "weight_decay": 0.0, "lr": new_lr}, ] return optim_groups def initialize_weights(self): self.apply(self._init_weights) # zero out classifier weights nn.init.zeros_(self.fc.weight) nn.init.zeros_(self.t_block[-1].weight) nn.init.zeros_(self.t_block[-1].bias) # zero out c_proj weights in all blocks for block in self.blocks: nn.init.zeros_(block.mlp.w3.weight) nn.init.zeros_(block.attn.proj.weight) # ControlNet-style identity init: zero the two paths that MEET at the # output -- block2.project (main) and to_out (skip) -- so the block emits # exactly 0 and stage 2 starts equivalent to stage 1. # # block1.project is deliberately left at normal init. Zeroing it too is a # trap: block1 out = 0 => silu(0) = 0 => grad(block2.project.weight) = 0, # and grad(block1) flows back through block2.project.weight = 0, so the two # lock each other at zero for the whole run. Both kernel-3 convs would stay # dead and the block would degenerate to to_out (a kernel-1 pointwise map) # plus a constant bias -- no temporal smearing, no nonlinearity. zero_init_local_embedder(self.local_embedder) def _init_weights(self, module): if isinstance(module, nn.Linear): # https://arxiv.org/pdf/2310.17813 fan_out = module.weight.size(0) fan_in = module.weight.size(1) std = 1.0 / math.sqrt(fan_in) * min(1.0, math.sqrt(fan_out / fan_in)) nn.init.normal_(module.weight, mean=0.0, std=std) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=0.02) def forward(self, x, t, bpm, rms, density, zcr, flatness, chroma, style, unconditional_mask=None): assert self.stage in [1, 2], f'Stage must be 1 or 2 but got {self.stage}' B, T, N, C = x.shape rms = (rms - self.rms_mean) / self.rms_std density = (density - self.density_mean) / self.density_std zcr = (zcr - self.zcr_mean) / self.zcr_std flatness = (flatness - self.flatness_mean) / self.flatness_std chroma = (chroma - self.chroma_mean) / self.chroma_std t = self.t_embedder(t.flatten()).view(B, T, -1) style = self.style_embedder(style) bpm = self.bpm_embedder(torch.clamp(torch.round(bpm), min=0, max=349).long()) # local (per-measure acoustic) signals: dropped to their normalized mean (0), # with a presence-mask channel carrying the "absent" information instead. local_keys = ['chroma', 'rms', 'density', 'zcr', 'flatness'] # (B, T) boolean drop masks, default = nothing dropped local_drop = {k: torch.zeros(B, T, dtype=torch.bool, device=x.device) for k in local_keys} if self.use_null_token: scalar_zero = torch.tensor(0.0, device=x.device, dtype=x.dtype) signals = { 'style': style, 'chroma': chroma, 'bpm': bpm, 'rms': rms, 'density': density, 'zcr': zcr, 'flatness': flatness, } null_tokens = { 'style': self.null_style, 'bpm': self.null_bpm, # local signals drop to their normalized mean; the presence channel (built below) is what encodes absence, so the null value is just 0. 'chroma': scalar_zero, 'rms': scalar_zero, 'density': scalar_zero, 'zcr': scalar_zero, 'flatness': scalar_zero, } if self.stage == 1: signals, drop_masks = multi_token_drop( signals, null_tokens, self.training, p_joint_uncond=0.1, p_joint_full=0.9, p_one_hot=0, p_ind_uncond=0, p_ind_low=0, p_ind_high=0, return_masks=True, ) elif self.stage == 2: # probabilities taken from Composer https://arxiv.org/pdf/2302.09778 signals, drop_masks = multi_token_drop( signals, null_tokens, self.training, p_joint_uncond=0.1, p_joint_full=0.1, p_one_hot=0, p_ind_uncond=0.5, p_ind_low=0, p_ind_high=0, return_masks=True, ) style = signals['style'] chroma = signals['chroma'] bpm = signals['bpm'] rms = signals['rms'] density = signals['density'] zcr = signals['zcr'] flatness = signals['flatness'] for k in local_keys: local_drop[k] = drop_masks[k] if unconditional_mask is not None: # style/bpm live in embedding space -> swap in learned null vectors style = torch.where(unconditional_mask['style'], self.null_style, style) bpm = torch.where(unconditional_mask['bpm'], self.null_bpm, bpm) # local signals -> zero the value AND fold the request into the drop mask so the presence channel is turned off too. chroma = torch.where(unconditional_mask['chroma'], scalar_zero, chroma) rms = torch.where(unconditional_mask['rms'].squeeze(-1), scalar_zero, rms) density = torch.where(unconditional_mask['density'].squeeze(-1), scalar_zero, density) zcr = torch.where(unconditional_mask['zcr'].squeeze(-1), scalar_zero, zcr) flatness = torch.where(unconditional_mask['flatness'].squeeze(-1), scalar_zero, flatness) for k in local_keys: local_drop[k] = local_drop[k] | unconditional_mask[k].squeeze(-1) x = rearrange(x, 'b t n c -> (b t) c n') x = self.x_embedder(x) x = rearrange(x, '(b t) c n -> b t n c', b=B, t=T) if self.stage == 2: # presence channels: 1 where the signal is present, 0 where dropped. # chroma shares a single mask, each scalar gets its own. presence = torch.stack([ (~local_drop['chroma']).to(x.dtype), (~local_drop['rms']).to(x.dtype), (~local_drop['density']).to(x.dtype), (~local_drop['zcr']).to(x.dtype), (~local_drop['flatness']).to(x.dtype), ], dim=-1) # (B, T, 5) c = torch.cat([chroma, rms.unsqueeze(-1), density.unsqueeze(-1), zcr.unsqueeze(-1), flatness.unsqueeze(-1), presence], dim=-1) # (B, T, 21) # convolve over the measure axis T (stride 1, kernel 3) so each measure's # descriptors smear into its neighbours, then broadcast the per-measure # embedding across the N within-measure latent positions. c = rearrange(c, 'b t f -> b f t') c = self.local_embedder(c) c = rearrange(c, 'b h t -> b t h') x = x + c.unsqueeze(2) x = rearrange(x, 'b t n c -> b (t n) c', b=B, t=T) t = t + style + bpm t0 = self.t_block(t) freqs_cis = self.freqs_cis[:x.shape[1]] for block in self.blocks: if self.gradient_checkpointing and self.training: x = checkpoint(block, x, t0, freqs_cis=freqs_cis, use_reentrant=False) else: x = block(x, t0, freqs_cis=freqs_cis) # SAM Audio does not use a non-linearity on t here shift, scale = (self.final_layer_scale_shift_table[None, None] + F.silu(t[:, :, None])).chunk( 2, dim=2 ) x = rearrange(x, 'b (t n) c -> b t n c', t=T, n=N // self.patch_size) x = modulate(self.norm(x), shift.expand(-1, -1, N // self.patch_size, -1), scale.expand(-1, -1, N // self.patch_size, -1)) x = self.fc(x) + self.bias x = rearrange(x, 'b t n (p c) -> b t (n p) c', p=self.patch_size, c=C) return x class MetaConditionalModernDiTV2Wrapper(nn.Module): def __init__(self, **kwargs): super().__init__() self.net = MetaConditionalModernDiTV2(**kwargs) self.diffusion = FM(timescale=1000.0) self.sampler = FMEulerSampler(self.diffusion) def forward(self, x, bpm, rms, density, zcr, flatness, chroma, style, t=None): return self.diffusion.loss( self.net, x, t=t, net_kwargs={ 'style': style, 'chroma': chroma, 'bpm': bpm, 'rms': rms, 'density': density, 'zcr': zcr, 'flatness': flatness, } ) def generate(self, shape, net_kwargs=None, uncond_net_kwargs=None, n_steps=50, guidance=1.0, noise=None, memory_efficient=True, rescale_phi=0, cfg_mode="independent", t_dist="uniform"): return self.sampler.sample( self.net, shape, n_steps=n_steps, net_kwargs=net_kwargs, uncond_net_kwargs=uncond_net_kwargs, guidance=guidance, noise=noise, memory_efficient=memory_efficient, rescale_phi=rescale_phi, cfg_mode=cfg_mode, t_dist=t_dist ) class _GradientBalancerFunction(torch.autograd.Function): """ The hidden autograd engine that intercepts the gradient flowing from a specific head down into the shared trunk, scaling it on the fly. """ @staticmethod def forward(ctx, features, total_buffer, fix_buffer, task_weight, total_weight, ema_decay, total_norm, epsilon): # Save variables for the backward pass ctx.total_buffer = total_buffer ctx.fix_buffer = fix_buffer ctx.task_weight = task_weight ctx.total_weight = total_weight ctx.ema_decay = ema_decay ctx.total_norm = total_norm ctx.epsilon = epsilon # Pass features through untouched during the forward pass return features.clone() @staticmethod def backward(ctx, grad_output): # 1. Compute per-batch-item norm (EnCodec style) dims = tuple(range(1, grad_output.dim())) norm = grad_output.norm(dim=dims).mean() batch_size = grad_output.shape[0] # 2. Update EMA buffers IN-PLACE # (Using pure tensor ops: no float(), no .item() -> zero graph breaks!) ctx.total_buffer.mul_(ctx.ema_decay).add_(norm * batch_size) ctx.fix_buffer.mul_(ctx.ema_decay).add_(batch_size) # 3. Calculate the smoothed average norm avg_norm = ctx.total_buffer / ctx.fix_buffer # 4. Calculate the EnCodec scaling factor ratio = ctx.task_weight / ctx.total_weight scale = ratio * ctx.total_norm / (ctx.epsilon + avg_norm) # 5. Scale the gradient before it passes down into the trunk grad_input = grad_output * scale # Return gradients for the inputs (None for the hyperparameter arguments) return grad_input, None, None, None, None, None, None, None class GradientBalancer(nn.Module): """ A drop-in module that wraps the EnCodec balancing math into an automatic layer. """ def __init__(self, weights: dict, ema_decay=0.999, total_norm=1.0, epsilon=1e-12): super().__init__() self.weights = weights self.ema_decay = ema_decay self.total_norm = total_norm self.epsilon = epsilon self.total_weight = sum(weights.values()) # Register EMA trackers as PyTorch buffers so they live on the GPU # and are saved in your model's state_dict for task in weights.keys(): self.register_buffer(f'total_{task}', torch.tensor(0.0)) self.register_buffer(f'fix_{task}', torch.tensor(0.0)) def forward(self, features, task_name): """ Pass the trunk features through this function before sending them to a head. """ # Fetch the specific EMA buffers for this task total_buffer = getattr(self, f'total_{task_name}') fix_buffer = getattr(self, f'fix_{task_name}') task_weight = self.weights[task_name] return _GradientBalancerFunction.apply( features, total_buffer, fix_buffer, task_weight, self.total_weight, self.ema_decay, self.total_norm, self.epsilon ) class PerceiverTokenPooler(nn.Module): def __init__(self, d_model: int, nhead: int = 8, mlp_ratio: float = 4.0): super().__init__() self.d_model = d_model # 1. The Single Latent Query Token (Learned Parameter) # We initialize it as (1, 1, d_model) so it easily broadcasts across batches self.latents = nn.Parameter(torch.randn(d_model) / d_model ** 0.5) # 2. Multi-Head Cross-Attention Layer # batch_first=True expects input shapes to be (batch, seq_len, features) self.cross_attn = nn.MultiheadAttention( embed_dim=d_model, num_heads=nhead, dropout=0, batch_first=True ) # 3. Standard Post-Attention Processing (LayerNorm + FeedForward) self.norm1 = nn.LayerNorm(d_model) self.norm2 = nn.LayerNorm(d_model) self.norm3 = nn.LayerNorm(d_model) self.mlp = SwiGLUMlp(d_model, int(2 / 3 * mlp_ratio * d_model), bias=True) @torch.compiler.disable def forward(self, signals: list[torch.Tensor]) -> torch.Tensor: """ Args: signals: A list of Tensors, where each tensor represents a processed signal. Each tensor in the list must have shape (batch_size, seq_len_i, d_model). Note: seq_len_i can vary between different signals! Returns: pooled_token: A tensor of shape (batch_size, 1, d_model) representing the single fused token for your DiT. """ B, T, C = signals[0].shape M = len(signals) combined = torch.cat([s.unsqueeze(2) for s in signals], dim=2) kv_sequence = self.norm1(combined.view(B * T, M, C)) q = self.norm2(self.latents.unsqueeze(0).unsqueeze(0).expand(B * T, -1, -1)) attn_output, _ = self.cross_attn( query=q, key=kv_sequence, value=kv_sequence ) x = q + attn_output x = x + self.mlp(self.norm3(x)) x = x.reshape(B, T, C) return x class MetaConditionalModernDiTV2Composer(nn.Module): def __init__(self, in_channels, hidden_size, spatial_window, n_chunks, style_dim, n_text_tokens, text_dim=1024, num_heads=12, depth=12, mlp_ratio=4, gradient_checkpointing=False, use_null_token=False, patch_size=1, drop_path_rate=0.1, signal_dim = {}, weights = {}, **kwargs, ): super().__init__() self.spatial_window = spatial_window self.gradient_checkpointing = gradient_checkpointing self.n_chunks = n_chunks max_input_size = spatial_window * n_chunks + n_text_tokens self.patch_size = patch_size self.signal_dim = signal_dim self.use_null_token = use_null_token self.balancer = GradientBalancer(weights=weights) self.t_embedder = TimestepEmbedder(hidden_size, bias=False, swiglu=True) self.local_embedder = Patcher(16, hidden_size, patch_size=patch_size, bias=True) self.style_embedder = nn.Linear(style_dim, hidden_size, bias=True) self.bpm_embedder = nn.Linear(768, hidden_size, bias=True) self.text_embedder = nn.Sequential(nn.LayerNorm(text_dim), nn.Linear(text_dim, hidden_size, bias=True)) self.text_embedder = nn.Sequential( #nn.LayerNorm(text_dim), nn.Linear(text_dim, hidden_size, bias=True), nn.SiLU(), nn.Linear(hidden_size, hidden_size, bias=True) ) self.pooler = PerceiverTokenPooler(hidden_size, num_heads, mlp_ratio) self.text_embed = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.audio_embed = nn.Parameter(torch.randn(hidden_size) / hidden_size ** 0.5) self.t_block = nn.Sequential( nn.SiLU(), nn.Linear(hidden_size, hidden_size * 6, bias=True), ) dp_rates=[x.item() for x in torch.linspace(0, drop_path_rate, depth)] self.blocks = nn.ModuleList([ DiTAirBlock(hidden_size, num_heads, mlp_ratio=mlp_ratio, drop_path=dp_rates[i]) for i in range(depth) ]) self.norm = nn.ModuleDict({ name: RMSNorm(hidden_size) for name in signal_dim.keys() }) self.final_layer_scale_shift_table = nn.ParameterDict({ name: nn.Parameter(torch.randn(2, hidden_size) / hidden_size ** 0.5,) for name in signal_dim.keys() }) self.fc = nn.ModuleDict({ name: nn.Linear(hidden_size, dim * patch_size, bias=False) for name, dim in signal_dim.items() }) self.bias = nn.ParameterDict({ name: nn.Parameter(torch.zeros(dim * patch_size)) for name, dim in signal_dim.items() }) self.register_buffer('freqs_cis', precompute_freqs_cis(hidden_size // num_heads, max_input_size)) self.initialize_weights() def create_optimizer_groups(self, weight_decay=1e-2, lr=1e-4): decay = [] no_decay = [] for name, param in self.named_parameters(): if not param.requires_grad: continue if param.ndim < 2: no_decay.append(param) else: decay.append(param) optim_groups = [ {"params": decay, "weight_decay": weight_decay, "lr": lr}, {"params": no_decay, "weight_decay": 0.0, "lr": lr}, ] return optim_groups def initialize_weights(self): self.apply(self._init_weights) # zero out classifier weights for name in self.signal_dim.keys(): nn.init.zeros_(self.fc[name].weight) nn.init.zeros_(self.t_block[-1].weight) nn.init.zeros_(self.t_block[-1].bias) # zero out c_proj weights in all blocks for block in self.blocks: nn.init.zeros_(block.mlp.w3.weight) nn.init.zeros_(block.attn.proj.weight) def _init_weights(self, module): if isinstance(module, nn.Linear): # https://arxiv.org/pdf/2310.17813 fan_out = module.weight.size(0) fan_in = module.weight.size(1) std = 1.0 / math.sqrt(fan_in) * min(1.0, math.sqrt(fan_out / fan_in)) nn.init.normal_(module.weight, mean=0.0, std=std) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=0.02) def forward(self, x, t, text, unconditional_mask=None): x = x.squeeze(2) style = x[..., :128] chroma = x[..., 128:128+12] rms = x[..., [128+12]] density = x[..., [128+13]] zcr = x[..., [128+14]] flatness = x[..., [128+15]] bpm = x[..., 128+16:128+16+768] # text null token is embedding of empty string if self.use_null_token: signals = { 'text': text, } null_tokens = { 'text': self.null_text, } signals = multi_token_drop( signals, null_tokens, self.training, p_joint_uncond=0.1, p_joint_full=0.9, p_one_hot=0, p_ind_uncond=0, p_ind_low=0, p_ind_high=0 ) text = signals['text'] if unconditional_mask is not None: text = torch.where(unconditional_mask['text'], null_tokens['text'], text) # prepend x with text x = torch.cat([chroma, rms, density, zcr, flatness], dim=-1) x = self.local_embedder(x.transpose(1, 2)).transpose(1, 2) bpm = self.bpm_embedder(bpm) style = self.style_embedder(style) x = self.pooler([x, bpm, style]) + self.audio_embed text = self.text_embedder(text) + self.text_embed x = torch.cat([text, x], dim=1) B, T = t.shape t = self.t_embedder(t.flatten()).view(B, T, -1) t = t + text.mean(dim=1, keepdims=True) # mean pool text for global embedder t0 = self.t_block(t) freqs_cis = self.freqs_cis[:x.shape[1]] for block in self.blocks: if self.gradient_checkpointing and self.training: x = checkpoint(block, x, t0, freqs_cis=freqs_cis, use_reentrant=False) else: x = block(x, t0, freqs_cis=freqs_cis) x = x[:, -self.n_chunks:] out = {} for name in self.signal_dim.keys(): features = self.balancer(x, name) # SAM Audio does not use a non-linearity on t here shift, scale = (self.final_layer_scale_shift_table[name][None] + F.silu(t[:, :, None])).chunk( 2, dim=2 ) features = modulate(self.norm[name](features), shift.squeeze(2), scale.squeeze(2)) features = self.fc[name](features) + self.bias[name] out[name] = features out = torch.cat(list(out.values()), dim=-1).unsqueeze(2) return out class MetaConditionalModernDiTV2ComposerWrapper(nn.Module): def __init__(self, **kwargs): super().__init__() self.net = MetaConditionalModernDiTV2Composer(**kwargs) self.diffusion = FM(timescale=1000.0) self.sampler = FMEulerSampler(self.diffusion) def forward(self, x, text, t=None): return self.diffusion.loss( self.net, x, t=t, net_kwargs={'text': text} ) def generate(self, shape, net_kwargs=None, uncond_net_kwargs=None, n_steps=50, guidance=1.0, noise=None, memory_efficient=True, rescale_phi=0, cfg_mode="independent", t_dist="uniform"): return self.sampler.sample( self.net, shape, n_steps=n_steps, net_kwargs=net_kwargs, uncond_net_kwargs=uncond_net_kwargs, guidance=guidance, noise=noise, memory_efficient=memory_efficient, rescale_phi=rescale_phi, cfg_mode=cfg_mode, t_dist=t_dist ) def ModernDiT_large(**kwargs): return ModernDiTWrapper(depth=28, hidden_size=1152, num_heads=16, **kwargs) def ModernDiT_medium(**kwargs): return ModernDiTWrapper(depth=24, hidden_size=1024, num_heads=16, **kwargs) def ModernDiT_small(**kwargs): return ModernDiTWrapper(depth=16, hidden_size=1024, num_heads=16, **kwargs) def ModernDiT_tiny(**kwargs): return ModernDiTWrapper(depth=16, hidden_size=768, num_heads=12, **kwargs) def UnconditionalModernDiT_large(**kwargs): return UnconditionalModernDiTWrapper(depth=28, hidden_size=1152, num_heads=16, **kwargs) def UnconditionalModernDiT_medium(**kwargs): return UnconditionalModernDiTWrapper(depth=24, hidden_size=1024, num_heads=16, **kwargs) def UnconditionalModernDiT_smedium_W5(**kwargs): return UnconditionalModernDiTWrapper(depth=20, hidden_size=1024, num_heads=16, **kwargs) def UnconditionalModernDiT_smedium_W4(**kwargs): return UnconditionalModernDiTWrapper(depth=20, hidden_size=768, num_heads=12, **kwargs) def UnconditionalModernDiT_smedium_W3(**kwargs): return UnconditionalModernDiTWrapper(depth=20, hidden_size=512, num_heads=8, **kwargs) def UnconditionalModernDiT_smedium_W2(**kwargs): return UnconditionalModernDiTWrapper(depth=20, hidden_size=384, num_heads=6, **kwargs) def UnconditionalModernDiT_smedium_W1(**kwargs): return UnconditionalModernDiTWrapper(depth=20, hidden_size=256, num_heads=4, **kwargs) def UnconditionalModernDiT_smedium_W0(**kwargs): return UnconditionalModernDiTWrapper(depth=20, hidden_size=128, num_heads=2, **kwargs) def UnconditionalModernDiT_smedium_D5(**kwargs): return UnconditionalModernDiTWrapper(depth=28, hidden_size=768, num_heads=12, **kwargs) def UnconditionalModernDiT_smedium_D4(**kwargs): return UnconditionalModernDiTWrapper(depth=24, hidden_size=768, num_heads=12, **kwargs) def UnconditionalModernDiT_smedium_D3(**kwargs): return UnconditionalModernDiTWrapper(depth=20, hidden_size=768, num_heads=12, **kwargs) def UnconditionalModernDiT_smedium_D2(**kwargs): return UnconditionalModernDiTWrapper(depth=16, hidden_size=768, num_heads=12, **kwargs) def UnconditionalModernDiT_smedium_D1(**kwargs): return UnconditionalModernDiTWrapper(depth=12, hidden_size=768, num_heads=12, **kwargs) def UnconditionalModernDiT_smedium_D0(**kwargs): return UnconditionalModernDiTWrapper(depth=8, hidden_size=768, num_heads=12, **kwargs) def UnconditionalModernDiT_smedium(**kwargs): return UnconditionalModernDiTWrapper(depth=20, hidden_size=768, num_heads=12, **kwargs) def UnconditionalModernDiT_small(**kwargs): return UnconditionalModernDiTWrapper(depth=16, hidden_size=1024, num_heads=16, **kwargs) def UnconditionalModernDiT_tiny(**kwargs): return UnconditionalModernDiTWrapper(depth=16, hidden_size=768, num_heads=12, **kwargs) def StyleConditionalModernDiT_large(**kwargs): return StyleConditionalModernDiTWrapper(depth=28, hidden_size=1152, num_heads=16, **kwargs) def StyleConditionalModernDiT_medium(**kwargs): return StyleConditionalModernDiTWrapper(depth=24, hidden_size=1024, num_heads=16, **kwargs) def StyleConditionalModernDiT_smedium(**kwargs): return StyleConditionalModernDiTWrapper(depth=20, hidden_size=768, num_heads=12, **kwargs) def StyleConditionalModernDiT_small(**kwargs): return StyleConditionalModernDiTWrapper(depth=16, hidden_size=1024, num_heads=16, **kwargs) def StyleConditionalModernDiT_tiny(**kwargs): return StyleConditionalModernDiTWrapper(depth=16, hidden_size=768, num_heads=12, **kwargs) def BpmRmsChromaStyleConditionalModernDiT_large(**kwargs): return BpmRmsChromaStyleConditionalModernDiTWrapper(depth=28, hidden_size=1152, num_heads=16, **kwargs) def BpmRmsChromaStyleConditionalModernDiT_medium(**kwargs): return BpmRmsChromaStyleConditionalModernDiTWrapper(depth=24, hidden_size=1024, num_heads=16, **kwargs) def BpmRmsChromaStyleConditionalModernDiT_smedium(**kwargs): return BpmRmsChromaStyleConditionalModernDiTWrapper(depth=20, hidden_size=768, num_heads=12, **kwargs) def BpmRmsChromaStyleConditionalModernDiT_small(**kwargs): return BpmRmsChromaStyleConditionalModernDiTWrapper(depth=16, hidden_size=1024, num_heads=16, **kwargs) def BpmRmsChromaStyleConditionalModernDiT_tiny(**kwargs): return BpmRmsChromaStyleConditionalModernDiTWrapper(depth=16, hidden_size=768, num_heads=12, **kwargs) def MetaConditionalModernDiT_large(**kwargs): return MetaConditionalModernDiTWrapper(depth=28, hidden_size=1152, num_heads=16, **kwargs) def MetaConditionalModernDiT_medium(**kwargs): return MetaConditionalModernDiTWrapper(depth=24, hidden_size=1024, num_heads=16, **kwargs) def MetaConditionalModernDiT_smedium(**kwargs): return MetaConditionalModernDiTWrapper(depth=20, hidden_size=768, num_heads=12, **kwargs) def MetaConditionalModernDiT_small(**kwargs): return MetaConditionalModernDiTWrapper(depth=16, hidden_size=1024, num_heads=16, **kwargs) def MetaConditionalModernDiT_tiny(**kwargs): return MetaConditionalModernDiTWrapper(depth=16, hidden_size=768, num_heads=12, **kwargs) def MetaConditionalModernDiTV2_large(**kwargs): return MetaConditionalModernDiTV2Wrapper(depth=28, hidden_size=1152, num_heads=16, **kwargs) def MetaConditionalModernDiTV2_medium(**kwargs): return MetaConditionalModernDiTV2Wrapper(depth=24, hidden_size=1024, num_heads=16, **kwargs) def MetaConditionalModernDiTV2_smedium(**kwargs): return MetaConditionalModernDiTV2Wrapper(depth=20, hidden_size=768, num_heads=12, **kwargs) def MetaConditionalModernDiTV2_small(**kwargs): return MetaConditionalModernDiTV2Wrapper(depth=16, hidden_size=1024, num_heads=16, **kwargs) def MetaConditionalModernDiTV2_tiny(**kwargs): return MetaConditionalModernDiTV2Wrapper(depth=16, hidden_size=768, num_heads=12, **kwargs) def MetaConditionalModernDiTV2Composer_large(**kwargs): return MetaConditionalModernDiTV2ComposerWrapper(depth=28, hidden_size=1152, num_heads=16, **kwargs) def MetaConditionalModernDiTV2Composer_medium(**kwargs): return MetaConditionalModernDiTV2ComposerWrapper(depth=24, hidden_size=1024, num_heads=16, **kwargs) def MetaConditionalModernDiTV2Composer_smedium(**kwargs): return MetaConditionalModernDiTV2ComposerWrapper(depth=20, hidden_size=768, num_heads=12, **kwargs) def MetaConditionalModernDiTV2Composer_small(**kwargs): return MetaConditionalModernDiTV2ComposerWrapper(depth=16, hidden_size=1024, num_heads=16, **kwargs) def MetaConditionalModernDiTV2Composer_tiny(**kwargs): return MetaConditionalModernDiTV2ComposerWrapper(depth=16, hidden_size=768, num_heads=12, **kwargs)