Text-to-Audio
Transformers
Safetensors
qadit
feature-extraction
diffusion
dit
audio
educational
research
custom_code
Instructions to use QuarkML/QaDiT-160 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use QuarkML/QaDiT-160 with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-to-audio", model="QuarkML/QaDiT-160", trust_remote_code=True)# Load model directly from transformers import AutoModel model = AutoModel.from_pretrained("QuarkML/QaDiT-160", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
| """ | |
| QaDiT model for Hugging Face Transformers (`trust_remote_code=True`). | |
| Contains: | |
| * DiT backbone (ported from audio_dit/dit.py) | |
| * Cosine-schedule v-prediction DDIM sampler (from audio_dit/diffusion.py) | |
| * :class:`QaDiTModel` — ``PreTrainedModel`` with ``generate(prompt=...)`` | |
| Example:: | |
| from transformers import AutoModel | |
| model = AutoModel.from_pretrained("USER/qadit", trust_remote_code=True) | |
| model = model.to("cuda") | |
| out = model.generate("Rain falls on a metal roof with distant thunder") | |
| # out.audios: list[np.ndarray], mono float32 in [-1, 1] | |
| """ | |
| from __future__ import annotations | |
| import math | |
| from dataclasses import dataclass | |
| from typing import List, Optional, Union | |
| import numpy as np | |
| import torch | |
| import torch.nn as nn | |
| import torch.nn.functional as F | |
| from transformers import PreTrainedModel | |
| from transformers.modeling_outputs import ModelOutput | |
| from transformers.utils import logging | |
| try: | |
| from .configuration_qadit import QaDiTConfig | |
| except ImportError: # Hub dynamic module loads files as siblings | |
| from configuration_qadit import QaDiTConfig | |
| logger = logging.get_logger(__name__) | |
| # --------------------------------------------------------------------------- # | |
| # DiT building blocks (self-contained for Hub upload) # | |
| # --------------------------------------------------------------------------- # | |
| class TimestepEmbedder(nn.Module): | |
| def __init__(self, hidden_size: int, freq_dim: int = 256): | |
| super().__init__() | |
| self.freq_dim = freq_dim | |
| self.mlp = nn.Sequential( | |
| nn.Linear(freq_dim, hidden_size), | |
| nn.SiLU(), | |
| nn.Linear(hidden_size, hidden_size), | |
| ) | |
| def sinusoidal(t: torch.Tensor, dim: int, max_period: int = 10_000) -> torch.Tensor: | |
| half = dim // 2 | |
| freqs = torch.exp( | |
| -math.log(max_period) | |
| * torch.arange(half, dtype=torch.float32, device=t.device) | |
| / half | |
| ) | |
| args = t.float()[:, None] * freqs[None] | |
| return torch.cat([torch.cos(args), torch.sin(args)], dim=-1) | |
| def forward(self, t: torch.Tensor) -> torch.Tensor: | |
| # Sinusoidal features are always built in float32 for precision; cast | |
| # to the weight dtype so half-precision backbones work unchanged. | |
| freqs = self.sinusoidal(t, self.freq_dim) | |
| return self.mlp(freqs.to(self.mlp[0].weight.dtype)) | |
| def build_2d_sincos_pos_embed(dim: int, grid_t: int, grid_f: int) -> torch.Tensor: | |
| assert dim % 4 == 0 | |
| def axis_embed(positions: torch.Tensor, axis_dim: int) -> torch.Tensor: | |
| omega = torch.arange(axis_dim // 2, dtype=torch.float32) / (axis_dim // 2) | |
| omega = 1.0 / (10_000 ** omega) | |
| out = positions.float()[:, None] * omega[None] | |
| return torch.cat([torch.sin(out), torch.cos(out)], dim=-1) | |
| t_pos = torch.arange(grid_t).repeat_interleave(grid_f) | |
| f_pos = torch.arange(grid_f).repeat(grid_t) | |
| return torch.cat( | |
| [axis_embed(t_pos, dim // 2), axis_embed(f_pos, dim // 2)], dim=-1 | |
| ) | |
| class PatchEmbed(nn.Module): | |
| def __init__(self, in_channels: int, hidden_size: int, patch_size: int): | |
| super().__init__() | |
| self.patch_size = patch_size | |
| self.proj = nn.Conv2d( | |
| in_channels, hidden_size, kernel_size=patch_size, stride=patch_size | |
| ) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| x = self.proj(x) | |
| return x.flatten(2).transpose(1, 2) | |
| class SelfAttention(nn.Module): | |
| def __init__(self, dim: int, num_heads: int): | |
| super().__init__() | |
| assert dim % num_heads == 0 | |
| self.num_heads = num_heads | |
| self.head_dim = dim // num_heads | |
| self.qkv = nn.Linear(dim, dim * 3) | |
| self.out = nn.Linear(dim, dim) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| B, N, D = x.shape | |
| qkv = self.qkv(x).reshape(B, N, 3, self.num_heads, self.head_dim) | |
| q, k, v = qkv.permute(2, 0, 3, 1, 4) | |
| x = F.scaled_dot_product_attention(q, k, v) | |
| return self.out(x.transpose(1, 2).reshape(B, N, D)) | |
| class CrossAttention(nn.Module): | |
| def __init__(self, dim: int, num_heads: int): | |
| super().__init__() | |
| assert dim % num_heads == 0 | |
| self.num_heads = num_heads | |
| self.head_dim = dim // num_heads | |
| self.q = nn.Linear(dim, dim) | |
| self.kv = nn.Linear(dim, dim * 2) | |
| self.out = nn.Linear(dim, dim) | |
| def forward( | |
| self, | |
| x: torch.Tensor, | |
| ctx: torch.Tensor, | |
| ctx_mask: Optional[torch.Tensor], | |
| ) -> torch.Tensor: | |
| B, N, D = x.shape | |
| L = ctx.shape[1] | |
| q = self.q(x).reshape(B, N, self.num_heads, self.head_dim).transpose(1, 2) | |
| kv = self.kv(ctx).reshape(B, L, 2, self.num_heads, self.head_dim) | |
| k, v = kv.permute(2, 0, 3, 1, 4) | |
| attn_mask = None | |
| if ctx_mask is not None: | |
| attn_mask = torch.where(ctx_mask.bool(), 0.0, float("-inf")) | |
| attn_mask = attn_mask[:, None, None, :].to(q.dtype) | |
| x = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask) | |
| return self.out(x.transpose(1, 2).reshape(B, N, D)) | |
| class MLP(nn.Module): | |
| def __init__(self, dim: int, hidden: int): | |
| super().__init__() | |
| self.fc1 = nn.Linear(dim, hidden) | |
| self.fc2 = nn.Linear(hidden, dim) | |
| def forward(self, x: torch.Tensor) -> torch.Tensor: | |
| return self.fc2(F.gelu(self.fc1(x), approximate="tanh")) | |
| def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor: | |
| return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1) | |
| class DiTBlock(nn.Module): | |
| def __init__(self, dim: int, num_heads: int, mlp_ratio: float): | |
| super().__init__() | |
| self.norm1 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) | |
| self.attn = SelfAttention(dim, num_heads) | |
| self.norm_ctx = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) | |
| self.cross = CrossAttention(dim, num_heads) | |
| self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) | |
| self.mlp = MLP(dim, int(dim * mlp_ratio)) | |
| self.adaLN = nn.Sequential(nn.SiLU(), nn.Linear(dim, 6 * dim)) | |
| nn.init.zeros_(self.adaLN[-1].weight) | |
| nn.init.zeros_(self.adaLN[-1].bias) | |
| nn.init.zeros_(self.cross.out.weight) | |
| nn.init.zeros_(self.cross.out.bias) | |
| def forward( | |
| self, | |
| x: torch.Tensor, | |
| cond: torch.Tensor, | |
| ctx: torch.Tensor, | |
| ctx_mask: Optional[torch.Tensor], | |
| ) -> torch.Tensor: | |
| ( | |
| shift_sa, | |
| scale_sa, | |
| gate_sa, | |
| shift_mlp, | |
| scale_mlp, | |
| gate_mlp, | |
| ) = self.adaLN(cond).chunk(6, dim=-1) | |
| x = x + gate_sa.unsqueeze(1) * self.attn( | |
| modulate(self.norm1(x), shift_sa, scale_sa) | |
| ) | |
| x = x + self.cross(self.norm_ctx(x), ctx, ctx_mask) | |
| x = x + gate_mlp.unsqueeze(1) * self.mlp( | |
| modulate(self.norm2(x), shift_mlp, scale_mlp) | |
| ) | |
| return x | |
| class FinalLayer(nn.Module): | |
| def __init__(self, dim: int, patch_size: int, out_channels: int): | |
| super().__init__() | |
| self.norm = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6) | |
| self.linear = nn.Linear(dim, patch_size * patch_size * out_channels) | |
| self.adaLN = nn.Sequential(nn.SiLU(), nn.Linear(dim, 2 * dim)) | |
| nn.init.zeros_(self.adaLN[-1].weight) | |
| nn.init.zeros_(self.adaLN[-1].bias) | |
| nn.init.zeros_(self.linear.weight) | |
| nn.init.zeros_(self.linear.bias) | |
| def forward(self, x: torch.Tensor, cond: torch.Tensor) -> torch.Tensor: | |
| shift, scale = self.adaLN(cond).chunk(2, dim=-1) | |
| return self.linear(modulate(self.norm(x), shift, scale)) | |
| class DiT(nn.Module): | |
| """Text-conditioned Diffusion Transformer over VAE mel-latents.""" | |
| def __init__( | |
| self, | |
| latent_channels: int, | |
| latent_time: int, | |
| latent_freq: int, | |
| patch_size: int, | |
| hidden_size: int, | |
| depth: int, | |
| num_heads: int, | |
| mlp_ratio: float, | |
| text_dim: int, | |
| repa_layer: int, | |
| ): | |
| super().__init__() | |
| assert latent_time % patch_size == 0 and latent_freq % patch_size == 0 | |
| self.out_channels = latent_channels | |
| self.hidden_size = hidden_size | |
| self.patch_size = patch_size | |
| self.grid_t = latent_time // patch_size | |
| self.grid_f = latent_freq // patch_size | |
| self.num_tokens = self.grid_t * self.grid_f | |
| self.repa_layer = repa_layer | |
| self.patch_embed = PatchEmbed(latent_channels, hidden_size, patch_size) | |
| # Built on first use rather than registered as a buffer: from_pretrained | |
| # materializes the model on the meta device, so a derived buffer that no | |
| # checkpoint supplies would silently load as zeros and corrupt every | |
| # sample. Recomputing it is exact, cheap and cached per device/dtype. | |
| self._pos_embed_cache: Optional[torch.Tensor] = None | |
| self.t_embed = TimestepEmbedder(hidden_size) | |
| self.text_proj = nn.Sequential( | |
| nn.LayerNorm(text_dim), | |
| nn.Linear(text_dim, hidden_size), | |
| ) | |
| self.pooled_proj = nn.Sequential( | |
| nn.SiLU(), | |
| nn.Linear(hidden_size, hidden_size), | |
| ) | |
| self.null_text = nn.Parameter(torch.zeros(1, 1, hidden_size)) | |
| self.blocks = nn.ModuleList( | |
| DiTBlock(hidden_size, num_heads, mlp_ratio) for _ in range(depth) | |
| ) | |
| self.final = FinalLayer(hidden_size, patch_size, self.out_channels) | |
| self._init_weights() | |
| def _init_weights(self): | |
| def basic(m): | |
| if isinstance(m, nn.Linear): | |
| nn.init.xavier_uniform_(m.weight) | |
| if m.bias is not None: | |
| nn.init.zeros_(m.bias) | |
| self.apply(basic) | |
| w = self.patch_embed.proj.weight | |
| nn.init.xavier_uniform_(w.view(w.shape[0], -1)) | |
| nn.init.zeros_(self.patch_embed.proj.bias) | |
| nn.init.normal_(self.t_embed.mlp[0].weight, std=0.02) | |
| nn.init.normal_(self.t_embed.mlp[2].weight, std=0.02) | |
| for block in self.blocks: | |
| nn.init.zeros_(block.adaLN[-1].weight) | |
| nn.init.zeros_(block.adaLN[-1].bias) | |
| nn.init.zeros_(block.cross.out.weight) | |
| nn.init.zeros_(block.cross.out.bias) | |
| nn.init.zeros_(self.final.adaLN[-1].weight) | |
| nn.init.zeros_(self.final.adaLN[-1].bias) | |
| nn.init.zeros_(self.final.linear.weight) | |
| nn.init.zeros_(self.final.linear.bias) | |
| def pos_embed(self, device: torch.device, dtype: torch.dtype) -> torch.Tensor: | |
| cache = self._pos_embed_cache | |
| if cache is None or cache.device != device or cache.dtype != dtype: | |
| cache = build_2d_sincos_pos_embed( | |
| self.hidden_size, self.grid_t, self.grid_f | |
| )[None].to(device=device, dtype=dtype) | |
| self._pos_embed_cache = cache | |
| return cache | |
| def null_context(self, batch_size: int): | |
| ctx = self.null_text.expand(batch_size, -1, -1) | |
| mask = torch.ones(batch_size, 1, device=ctx.device, dtype=torch.long) | |
| return ctx, mask | |
| def unpatchify(self, x: torch.Tensor) -> torch.Tensor: | |
| B = x.shape[0] | |
| p, c = self.patch_size, self.out_channels | |
| x = x.reshape(B, self.grid_t, self.grid_f, p, p, c) | |
| x = torch.einsum("btfpqc->bctpfq", x) | |
| return x.reshape(B, c, self.grid_t * p, self.grid_f * p) | |
| def forward( | |
| self, | |
| z_t: torch.Tensor, | |
| t: torch.Tensor, | |
| text_emb: Optional[torch.Tensor], | |
| text_mask: Optional[torch.Tensor], | |
| drop_mask: Optional[torch.Tensor] = None, | |
| return_repa_hidden: bool = False, | |
| ): | |
| B = z_t.shape[0] | |
| x = self.patch_embed(z_t) | |
| x = x + self.pos_embed(x.device, x.dtype) | |
| if text_emb is None: | |
| ctx, ctx_mask = self.null_context(B) | |
| else: | |
| ctx = self.text_proj(text_emb) | |
| ctx_mask = text_mask | |
| if drop_mask is not None: | |
| null = self.null_text.expand(B, ctx.shape[1], -1) | |
| ctx = torch.where(drop_mask[:, None, None], null, ctx) | |
| null_mask = torch.zeros_like(ctx_mask) | |
| null_mask[:, 0] = 1 | |
| ctx_mask = torch.where(drop_mask[:, None], null_mask, ctx_mask) | |
| cond = self.t_embed(t) | |
| if ctx_mask is not None: | |
| denom = ctx_mask.sum(dim=1, keepdim=True).clamp(min=1) | |
| pooled = (ctx * ctx_mask.unsqueeze(-1)).sum(dim=1) / denom | |
| else: | |
| pooled = ctx.mean(dim=1) | |
| cond = cond + self.pooled_proj(pooled) | |
| repa_hidden = None | |
| for i, block in enumerate(self.blocks): | |
| x = block(x, cond, ctx, ctx_mask) | |
| if return_repa_hidden and i == self.repa_layer: | |
| repa_hidden = x | |
| out = self.unpatchify(self.final(x, cond)) | |
| if return_repa_hidden: | |
| return out, repa_hidden | |
| return out | |
| # --------------------------------------------------------------------------- # | |
| # Diffusion schedule + DDIM # | |
| # --------------------------------------------------------------------------- # | |
| class DiffusionScheduler: | |
| """Cosine alpha-bar schedule with v-prediction DDIM + CFG.""" | |
| def __init__( | |
| self, | |
| num_train_steps: int = 1000, | |
| schedule: str = "cosine", | |
| logit_normal_mean: float = 0.0, | |
| logit_normal_std: float = 1.0, | |
| ): | |
| self.T = num_train_steps | |
| self.ln_mean = logit_normal_mean | |
| self.ln_std = logit_normal_std | |
| self.schedule = schedule | |
| if schedule != "cosine": | |
| raise ValueError(f"unknown schedule: {schedule}") | |
| # May be created on the meta device under HF's init_empty_weights(); | |
| # materialize_real() / to() rebuilds a real CPU/CUDA table before use. | |
| self.alpha_bar = self._build_alpha_bar(self.T) | |
| def _build_alpha_bar(num_train_steps: int, device=None) -> torch.Tensor: | |
| device = torch.device(device) if device is not None else torch.device("cpu") | |
| # Force a concrete device — never allocate on "meta". | |
| if device.type == "meta": | |
| device = torch.device("cpu") | |
| s = 0.008 | |
| steps = torch.arange( | |
| num_train_steps + 1, dtype=torch.float64, device=device | |
| ) | |
| f = torch.cos((steps / num_train_steps + s) / (1 + s) * math.pi / 2) ** 2 | |
| abar = (f / f[0]).clamp(1e-5, 1.0) | |
| return abar[1:].float() | |
| def _is_meta(self) -> bool: | |
| t = self.alpha_bar | |
| return bool(getattr(t, "is_meta", False) or t.device.type == "meta") | |
| def materialize_real(self, device=None) -> "DiffusionScheduler": | |
| """Rebuild alpha_bar if it was left on the meta device by from_pretrained.""" | |
| target = torch.device(device) if device is not None else torch.device("cpu") | |
| if target.type == "meta": | |
| target = torch.device("cpu") | |
| if self._is_meta() or self.alpha_bar.device != target: | |
| self.alpha_bar = self._build_alpha_bar(self.T, device="cpu").to(target) | |
| return self | |
| def to(self, device) -> "DiffusionScheduler": | |
| return self.materialize_real(device) | |
| def _gather(self, t: torch.Tensor): | |
| if self._is_meta(): | |
| self.materialize_real(t.device) | |
| abar = self.alpha_bar.to(t.device)[t] | |
| return abar.sqrt().view(-1, 1, 1, 1), (1 - abar).sqrt().view(-1, 1, 1, 1) | |
| def z0_from_v(self, z_t, t, v): | |
| sqrt_abar, sqrt_1m = self._gather(t) | |
| return sqrt_abar * z_t - sqrt_1m * v | |
| def eps_from_v(self, z_t, t, v): | |
| sqrt_abar, sqrt_1m = self._gather(t) | |
| return sqrt_1m * z_t + sqrt_abar * v | |
| def ddim_sample( | |
| self, | |
| model: nn.Module, | |
| shape: tuple, | |
| text_emb: torch.Tensor, | |
| text_mask: torch.Tensor, | |
| num_steps: int = 50, | |
| guidance_scale: float = 4.0, | |
| eta: float = 0.0, | |
| device: Union[str, torch.device] = "cpu", | |
| generator: Optional[torch.Generator] = None, | |
| dtype: Optional[torch.dtype] = None, | |
| ) -> torch.Tensor: | |
| self.materialize_real(device) | |
| B = shape[0] | |
| z = torch.randn(shape, device=device, generator=generator) | |
| times = torch.linspace(self.T - 1, 0, num_steps, device=device).long() | |
| use_cfg = guidance_scale is not None and guidance_scale > 1.0 | |
| for i in range(num_steps): | |
| t = times[i].expand(B) | |
| # Keep the schedule arithmetic in float32 even when the backbone | |
| # runs in half precision: the DDIM update is sensitive to it. | |
| z_in = z.to(dtype) if dtype is not None else z | |
| if use_cfg: | |
| v_cond = model(z_in, t.float(), text_emb, text_mask) | |
| v_uncond = model(z_in, t.float(), None, None) | |
| v = v_uncond + guidance_scale * (v_cond - v_uncond) | |
| else: | |
| v = model(z_in, t.float(), text_emb, text_mask) | |
| v = v.float() | |
| z0_hat = self.z0_from_v(z, t, v) | |
| eps_hat = self.eps_from_v(z, t, v) | |
| if i == num_steps - 1: | |
| z = z0_hat | |
| break | |
| t_next = times[i + 1].expand(B) | |
| abar_next = self.alpha_bar[t_next].view(-1, 1, 1, 1) | |
| abar_now = self.alpha_bar[t].view(-1, 1, 1, 1) | |
| sigma = eta * torch.sqrt( | |
| (1 - abar_next) / (1 - abar_now) * (1 - abar_now / abar_next) | |
| ) | |
| noise = ( | |
| torch.randn(shape, device=device, generator=generator) | |
| if eta > 0 | |
| else torch.zeros_like(z) | |
| ) | |
| dir_zt = torch.sqrt((1 - abar_next - sigma ** 2).clamp(min=0.0)) * eps_hat | |
| z = abar_next.sqrt() * z0_hat + dir_zt + sigma * noise | |
| return z | |
| # --------------------------------------------------------------------------- # | |
| # HF model outputs # | |
| # --------------------------------------------------------------------------- # | |
| class QaDiTOutput(ModelOutput): | |
| """Output of :meth:`QaDiTModel.forward` (single denoising step).""" | |
| sample: torch.FloatTensor = None | |
| class QaDiTGeneratorOutput(ModelOutput): | |
| """Output of :meth:`QaDiTModel.generate`. | |
| Hugging Face ``ModelOutput`` requires every field after the first to default | |
| to ``None`` (not other sentinels like ``16000``). | |
| """ | |
| audios: Optional[List[np.ndarray]] = None | |
| audio_values: Optional[torch.FloatTensor] = None | |
| latents: Optional[torch.FloatTensor] = None | |
| sampling_rate: Optional[int] = None | |
| # --------------------------------------------------------------------------- # | |
| # PreTrainedModel # | |
| # --------------------------------------------------------------------------- # | |
| class QaDiTModel(PreTrainedModel): | |
| """QaDiT: latent Diffusion Transformer for text-to-audio (~160M). | |
| Load with:: | |
| model = AutoModel.from_pretrained("USER/qadit", trust_remote_code=True) | |
| Then:: | |
| out = model.generate("A dog barks while birds chirp in the distance") | |
| # out.audios[0] -> np.ndarray, shape [num_samples], float32 | |
| """ | |
| config_class = QaDiTConfig | |
| base_model_prefix = "transformer" | |
| main_input_name = "latents" | |
| supports_gradient_checkpointing = False | |
| _no_split_modules = ["DiTBlock"] | |
| def __init__(self, config: QaDiTConfig): | |
| super().__init__(config) | |
| self.config = config | |
| self.transformer = DiT( | |
| latent_channels=config.latent_channels, | |
| latent_time=config.latent_time, | |
| latent_freq=config.latent_freq, | |
| patch_size=config.patch_size, | |
| hidden_size=config.hidden_size, | |
| depth=config.depth, | |
| num_heads=config.num_heads, | |
| mlp_ratio=config.mlp_ratio, | |
| text_dim=config.text_dim, | |
| repa_layer=config.repa_layer, | |
| ) | |
| self.scheduler = DiffusionScheduler( | |
| num_train_steps=config.num_train_timesteps, | |
| schedule=config.schedule, | |
| logit_normal_mean=config.logit_normal_mean, | |
| logit_normal_std=config.logit_normal_std, | |
| ) | |
| # Lazily populated by prepare_auxiliaries() / generate() | |
| self.tokenizer = None | |
| self.text_encoder = None | |
| self.vae = None | |
| self.vocoder = None | |
| self._aux_loaded = False | |
| self.post_init() | |
| # ------------------------------------------------------------------ # | |
| # Core forward (one denoising step) # | |
| # ------------------------------------------------------------------ # | |
| def forward( | |
| self, | |
| latents: torch.FloatTensor, | |
| timesteps: torch.FloatTensor, | |
| encoder_hidden_states: Optional[torch.FloatTensor] = None, | |
| encoder_attention_mask: Optional[torch.Tensor] = None, | |
| return_dict: bool = True, | |
| ): | |
| """Predict v for noisy ``latents`` at ``timesteps``. | |
| Parameters | |
| ---------- | |
| latents: | |
| ``[B, C, T, F]`` noisy latents in *scaled* training space. | |
| timesteps: | |
| ``[B]`` diffusion timesteps (float). | |
| encoder_hidden_states: | |
| T5 hidden states ``[B, L, text_dim]``, or ``None`` for unconditional. | |
| encoder_attention_mask: | |
| ``[B, L]`` with 1 = real token. | |
| """ | |
| sample = self.transformer( | |
| latents, | |
| timesteps, | |
| encoder_hidden_states, | |
| encoder_attention_mask, | |
| ) | |
| if not return_dict: | |
| return (sample,) | |
| return QaDiTOutput(sample=sample) | |
| # ------------------------------------------------------------------ # | |
| # Auxiliaries (T5 / VAE / vocoder) # | |
| # ------------------------------------------------------------------ # | |
| def prepare_auxiliaries(self, device: Optional[torch.device] = None): | |
| """Load frozen T5, AudioLDM VAE and HiFi-GAN if not already loaded.""" | |
| if self._aux_loaded: | |
| return self | |
| device = device or self.device | |
| cfg = self.config | |
| from transformers import AutoTokenizer, SpeechT5HifiGan, T5EncoderModel | |
| try: | |
| from diffusers import AutoencoderKL | |
| except ImportError as exc: | |
| raise ImportError( | |
| "diffusers is required for QaDiT waveform generation. " | |
| "Install with: pip install diffusers" | |
| ) from exc | |
| logger.info("Loading text encoder %s", cfg.text_model) | |
| self.tokenizer = AutoTokenizer.from_pretrained(cfg.text_model) | |
| self.text_encoder = ( | |
| T5EncoderModel.from_pretrained(cfg.text_model).to(device).eval() | |
| ) | |
| for p in self.text_encoder.parameters(): | |
| p.requires_grad_(False) | |
| logger.info("Loading VAE %s/%s", cfg.vae_model, cfg.vae_subfolder) | |
| self.vae = ( | |
| AutoencoderKL.from_pretrained(cfg.vae_model, subfolder=cfg.vae_subfolder) | |
| .to(device) | |
| .eval() | |
| ) | |
| for p in self.vae.parameters(): | |
| p.requires_grad_(False) | |
| logger.info( | |
| "Loading vocoder %s/%s", cfg.vocoder_model, cfg.vocoder_subfolder | |
| ) | |
| self.vocoder = ( | |
| SpeechT5HifiGan.from_pretrained( | |
| cfg.vocoder_model, subfolder=cfg.vocoder_subfolder | |
| ) | |
| .to(device) | |
| .eval() | |
| ) | |
| for p in self.vocoder.parameters(): | |
| p.requires_grad_(False) | |
| self._aux_loaded = True | |
| return self | |
| def encode_prompt( | |
| self, | |
| prompt: Union[str, List[str]], | |
| device: Optional[torch.device] = None, | |
| ): | |
| """Tokenize + T5-encode captions → ``(text_emb, text_mask)``.""" | |
| if not self._aux_loaded: | |
| self.prepare_auxiliaries(device) | |
| device = device or self.device | |
| if isinstance(prompt, str): | |
| prompt = [prompt] | |
| tok = self.tokenizer( | |
| prompt, | |
| padding="max_length", | |
| truncation=True, | |
| max_length=self.config.text_max_length, | |
| return_tensors="pt", | |
| ) | |
| input_ids = tok.input_ids.to(device) | |
| attention_mask = tok.attention_mask.to(device) | |
| with torch.no_grad(): | |
| text_emb = self.text_encoder( | |
| input_ids=input_ids, attention_mask=attention_mask | |
| ).last_hidden_state | |
| return text_emb, attention_mask | |
| # ------------------------------------------------------------------ # | |
| # Generation # | |
| # ------------------------------------------------------------------ # | |
| def generate( | |
| self, | |
| prompt: Optional[Union[str, List[str]]] = None, | |
| encoder_hidden_states: Optional[torch.FloatTensor] = None, | |
| encoder_attention_mask: Optional[torch.Tensor] = None, | |
| num_inference_steps: Optional[int] = None, | |
| guidance_scale: Optional[float] = None, | |
| seed: Optional[int] = 0, | |
| generator: Optional[Union[torch.Generator, List[torch.Generator]]] = None, | |
| eta: float = 0.0, | |
| output_type: str = "np", | |
| return_dict: bool = True, | |
| **kwargs, | |
| ): | |
| """Generate audio from text prompts. | |
| Parameters | |
| ---------- | |
| prompt: | |
| Caption string or list of captions. Ignored if | |
| ``encoder_hidden_states`` is provided. | |
| num_inference_steps: | |
| DDIM steps (default from config). | |
| guidance_scale: | |
| Classifier-free guidance scale (default from config). | |
| seed: | |
| Random seed used when ``generator`` is not supplied. Defaults to | |
| 0, matching the original ``audio_dit/sample.py`` CLI. | |
| output_type: | |
| ``"np"`` → numpy waveforms, ``"pt"`` → torch waveforms, | |
| ``"latent"`` → scaled latents only (no VAE/vocoder). | |
| """ | |
| cfg = self.config | |
| device = self.device | |
| # from_pretrained() normally returns eval mode, but make generation | |
| # invariant to callers having toggled train() in the same process. | |
| self.eval() | |
| steps = num_inference_steps or cfg.num_inference_steps | |
| guidance = ( | |
| guidance_scale if guidance_scale is not None else cfg.guidance_scale | |
| ) | |
| if cfg.latent_scale <= 0: | |
| raise ValueError( | |
| f"config.latent_scale must be positive, got {cfg.latent_scale}" | |
| ) | |
| if encoder_hidden_states is None: | |
| if prompt is None: | |
| raise ValueError("Provide `prompt` or `encoder_hidden_states`") | |
| if cfg.load_auxiliaries or output_type != "latent": | |
| self.prepare_auxiliaries(device) | |
| encoder_hidden_states, encoder_attention_mask = self.encode_prompt( | |
| prompt, device=device | |
| ) | |
| else: | |
| encoder_hidden_states = encoder_hidden_states.to(device) | |
| if encoder_attention_mask is not None: | |
| encoder_attention_mask = encoder_attention_mask.to(device) | |
| if isinstance(prompt, str): | |
| batch = 1 | |
| elif isinstance(prompt, list): | |
| batch = len(prompt) | |
| else: | |
| batch = encoder_hidden_states.shape[0] | |
| # silence unused | |
| _ = batch | |
| B = encoder_hidden_states.shape[0] | |
| shape = ( | |
| B, | |
| cfg.latent_channels, | |
| cfg.latent_time, | |
| cfg.latent_freq, | |
| ) | |
| if generator is None and seed is not None: | |
| generator = torch.Generator(device=device.type).manual_seed(seed) | |
| if isinstance(generator, list): | |
| if len(generator) != B: | |
| raise ValueError( | |
| f"Got {len(generator)} generators for batch size {B}" | |
| ) | |
| # Fall back to first generator for the shared noise draw; per-sample | |
| # generators are uncommon for this model. | |
| generator = generator[0] | |
| self.scheduler.to(device) | |
| latents = self.scheduler.ddim_sample( | |
| model=self.transformer, | |
| shape=shape, | |
| text_emb=encoder_hidden_states.to(self.dtype), | |
| text_mask=encoder_attention_mask, | |
| num_steps=steps, | |
| guidance_scale=guidance, | |
| eta=eta, | |
| device=device, | |
| generator=generator, | |
| dtype=self.dtype, | |
| ) | |
| if output_type == "latent": | |
| if not return_dict: | |
| return (latents,) | |
| return QaDiTGeneratorOutput( | |
| latents=latents, sampling_rate=cfg.sample_rate | |
| ) | |
| if not self._aux_loaded: | |
| self.prepare_auxiliaries(device) | |
| # Undo training latent scale, then VAE decode → mel → waveform. | |
| z = (latents / cfg.latent_scale).to(self.vae.dtype) | |
| mel = self.vae.decode(z).sample # [B, 1, 1024, 64] | |
| wav = self.vocoder(mel.squeeze(1).to(self.vocoder.dtype)) # [B, num_samples] | |
| wav = wav.float().clamp(-1, 1) | |
| if output_type == "pt": | |
| if not return_dict: | |
| return (wav, latents) | |
| return QaDiTGeneratorOutput( | |
| audio_values=wav, | |
| latents=latents, | |
| sampling_rate=cfg.sample_rate, | |
| ) | |
| # default: numpy | |
| audios = [w.detach().cpu().float().numpy() for w in wav] | |
| if not return_dict: | |
| return (audios, latents) | |
| return QaDiTGeneratorOutput( | |
| audios=audios, | |
| latents=latents, | |
| sampling_rate=cfg.sample_rate, | |
| ) | |
| def _set_gradient_checkpointing(self, module, value=False): | |
| pass | |
| # Register for Auto* when used as a local package / after from_pretrained | |
| try: | |
| QaDiTConfig.register_for_auto_class() | |
| QaDiTModel.register_for_auto_class("AutoModel") | |
| except Exception: | |
| # Older transformers or already-registered; auto_map in config.json still works. | |
| pass | |
| __all__ = [ | |
| "DiT", | |
| "DiffusionScheduler", | |
| "QaDiTModel", | |
| "QaDiTOutput", | |
| "QaDiTGeneratorOutput", | |
| ] | |