Skip to content

vllm_omni.diffusion.models.auk.auk_transformer

The AuK audio DiT as a self-contained inference module.

Two phases. num_layers double-stream blocks attend jointly over text and audio while keeping each stream's residual separate, then num_single_layers single-stream blocks run over the concatenated [text | ref audio | target audio] sequence. Only the target-audio slice is projected back to latent space, so the module predicts a flow-matching velocity for the target frames and :func:sample_latents integrates it.

Parameter names and shapes match the AuK reference backbone, so a released checkpoint loads with load_state_dict(strict=True) once the transformer. prefix is stripped (see :func:dit_state_dict).

Deliberate differences from the reference, all inference-only:

  • x_transformers and torchdiffeq are gone. Rotary embeddings are computed here (interleaved GPT-J pairs, no xpos scaling) and the ODE is an explicit Euler loop.
  • Dropout, activation checkpointing, the flash_attn backend switch and the zero-init of the modulation layers are training concerns and are dropped.
  • The timestep sinusoid is cast to the time MLP's weight dtype rather than to the timestep's own dtype, so an fp32 timestep works against half-precision weights outside autocast.
  • attn_mask_enabled defaults to True, the value the released config sets, rather than the reference's False.

AdaLayerNorm

Bases: Module

Timestep-conditioned modulation for a block: six chunks from one projection.

linear instance-attribute

linear = nn.Linear(dim, dim * 6)

norm instance-attribute

norm = nn.LayerNorm(
    dim, elementwise_affine=False, eps=1e-06
)

silu instance-attribute

silu = nn.SiLU()

forward

forward(
    x: Tensor, emb: Tensor
) -> tuple[Tensor, Tensor, Tensor, Tensor, Tensor]

AdaLayerNormFinal

Bases: Module

Timestep-conditioned modulation before the output projection.

linear instance-attribute

linear = nn.Linear(dim, dim * 2)

norm instance-attribute

norm = nn.LayerNorm(
    dim, elementwise_affine=False, eps=1e-06
)

silu instance-attribute

silu = nn.SiLU()

forward

forward(x: Tensor, emb: Tensor) -> Tensor

Attention

Bases: Module

Self-attention with per-head QK RMSNorm and rotary positions.

to_out stays a ModuleList because the released checkpoints name the output projection to_out.0; the reference's trailing dropout carries no parameters and is dropped.

attn_mask_enabled instance-attribute

attn_mask_enabled = attn_mask_enabled

dim_head instance-attribute

dim_head = dim_head

heads instance-attribute

heads = heads

k_norm instance-attribute

k_norm = nn.RMSNorm(dim_head, elementwise_affine=True)

q_norm instance-attribute

q_norm = nn.RMSNorm(dim_head, elementwise_affine=True)

to_out instance-attribute

to_out = nn.ModuleList([nn.Linear(inner_dim, dim)])

to_qkv instance-attribute

to_qkv = nn.Linear(dim, 3 * inner_dim)

forward

forward(
    x: Tensor,
    mask: Tensor | None = None,
    rope: Tensor | None = None,
) -> Tensor

AuKTransformer

Bases: Module

Flow-matching velocity predictor for AuK audio latents.

Parameters:

Name Type Description Default
dim int

Model width.

required
heads int

Attention heads.

required
dim_head int

Per-head width.

required
ff_mult float

Feed-forward expansion factor.

required
latent_dim int

Audio VAE latent channels, the input and output width.

required
text_hidden_dim int

Width of the pre-encoded text hidden states.

required
num_layers int

Double-stream (joint attention) block count.

required
num_single_layers int

Single-stream block count.

required
attn_mask_enabled bool

Apply padding masks inside attention. When False attention runs unmasked, though padded outputs are still zeroed, which is what the reference does with the flag off.

True

audio_embed instance-attribute

audio_embed = AudioEmbedding(latent_dim, dim)

dim instance-attribute

dim = dim

latent_dim instance-attribute

latent_dim = latent_dim

norm_out instance-attribute

norm_out = AdaLayerNormFinal(dim)

proj_out instance-attribute

proj_out = nn.Linear(dim, latent_dim)

rotary_embed instance-attribute

rotary_embed = Rotary(dim_head)

single_transformer_blocks instance-attribute

single_transformer_blocks = nn.ModuleList(
    [
        SingleBlock(
            dim=dim,
            heads=heads,
            dim_head=dim_head,
            ff_mult=ff_mult,
            attn_mask_enabled=attn_mask_enabled,
        )
        for _ in range(num_single_layers)
    ]
)

text_cond instance-attribute

text_cond: Tensor | None = None

text_uncond instance-attribute

text_uncond: Tensor | None = None

time_embed instance-attribute

time_embed = TimeEmbedding(dim)

transformer_blocks instance-attribute

transformer_blocks = nn.ModuleList(
    [
        DoubleBlock(
            dim=dim,
            heads=heads,
            dim_head=dim_head,
            ff_mult=ff_mult,
            attn_mask_enabled=attn_mask_enabled,
        )
        for _ in range(num_layers)
    ]
)

txt_norm instance-attribute

txt_norm = nn.RMSNorm(dim, elementwise_affine=True)

txt_proj instance-attribute

txt_proj = nn.Linear(text_hidden_dim, dim)

clear_cache

clear_cache() -> None

Drop the cached text projection. Call this when the conditioning changes.

forward

forward(
    x: Tensor,
    text: Tensor,
    time: Tensor,
    *,
    mask: Tensor | None = None,
    c_mask: Tensor | None = None,
    ref: Tensor | None = None,
    ref_mask: Tensor | None = None,
    drop_audio_cond: bool = False,
    drop_text: bool = False,
    cfg_infer: bool = False,
    cache: bool = False,
) -> Tensor

Predict the flow-matching velocity for the target frames of x.

Parameters:

Name Type Description Default
x Tensor

Noised target latents [B, n, latent_dim].

required
text Tensor

Pre-encoded text hidden states [B, nt, text_hidden_dim].

required
time Tensor

Flow-matching timestep, a scalar or [B].

required
mask Tensor | None

Target padding mask [B, n], True where valid.

None
c_mask Tensor | None

Text padding mask [B, nt]. Inferred from all-zero text rows when omitted.

None
ref Tensor | None

Reference prompt latents [B, np, latent_dim], prepended to the audio sequence. A zero-length ref counts as absent.

None
ref_mask Tensor | None

Reference padding mask [B, np].

None
drop_audio_cond bool

Zero the reference prompt, for the CFG uncond branch.

False
drop_text bool

Zero the projected text, for the CFG uncond branch.

False
cfg_infer bool

Run the cond and uncond branches as one batch of 2B, cond first. drop_audio_cond and drop_text are then driven per branch and ignored.

False
cache bool

Reuse the projected text across calls, for an ODE loop over fixed conditioning. Only the cfg_infer path caches, as in the reference. Call :meth:clear_cache when the text changes.

False

Returns:

Type Description
Tensor

Velocity [B, n, latent_dim], or [2B, n, latent_dim] under

Tensor

cfg_infer.

project_text

project_text(text: Tensor) -> Tensor

Project LLM hidden states [B, nt, text_hidden_dim] to model width.

AudioEmbedding

Bases: Module

Project audio latents to model width and add conv position information.

conv_pos_embed instance-attribute

conv_pos_embed = ConvPosEmbedding(dim)

linear instance-attribute

linear = nn.Linear(latent_dim, dim)

forward

forward(x: Tensor, mask: Tensor | None = None) -> Tensor

ConvPosEmbedding

Bases: Module

Depthwise-grouped conv stack that adds local position information.

conv1d instance-attribute

conv1d = nn.Sequential(
    nn.Conv1d(
        dim,
        dim,
        kernel_size,
        groups=groups,
        padding=padding,
    ),
    nn.Mish(),
    nn.Conv1d(
        dim,
        dim,
        kernel_size,
        groups=groups,
        padding=padding,
    ),
    nn.Mish(),
)

forward

forward(x: Tensor, mask: Tensor | None = None) -> Tensor

DoubleBlock

Bases: Module

MM-DiT block: joint attention, separate modulation and feed-forward per stream.

attn instance-attribute

attn = JointAttention(
    dim=dim,
    heads=heads,
    dim_head=dim_head,
    context_dim=dim,
    attn_mask_enabled=attn_mask_enabled,
)

attn_norm_c instance-attribute

attn_norm_c = AdaLayerNorm(dim)

attn_norm_x instance-attribute

attn_norm_x = AdaLayerNorm(dim)

ff_c instance-attribute

ff_c = FeedForward(dim=dim, mult=ff_mult)

ff_norm_c instance-attribute

ff_norm_c = nn.LayerNorm(
    dim, elementwise_affine=False, eps=1e-06
)

ff_norm_x instance-attribute

ff_norm_x = nn.LayerNorm(
    dim, elementwise_affine=False, eps=1e-06
)

ff_x instance-attribute

ff_x = FeedForward(dim=dim, mult=ff_mult)

forward

forward(
    x: Tensor,
    c: Tensor,
    t: Tensor,
    mask: Tensor | None,
    rope: Tensor,
    c_rope: Tensor,
    c_mask: Tensor | None,
) -> tuple[Tensor, Tensor]

FeedForward

Bases: Module

SwiGLU feed-forward with the gate projection fused into linear_in.

linear_in instance-attribute

linear_in = nn.Linear(dim, inner_dim * 2, bias=False)

linear_out instance-attribute

linear_out = nn.Linear(inner_dim, dim, bias=False)

forward

forward(x: Tensor) -> Tensor

JointAttention

Bases: Attention

Attention over the audio and text streams jointly, with separate projections.

c_k_norm instance-attribute

c_k_norm = nn.RMSNorm(dim_head, elementwise_affine=True)

c_q_norm instance-attribute

c_q_norm = nn.RMSNorm(dim_head, elementwise_affine=True)

to_out_c instance-attribute

to_out_c = nn.Linear(inner_dim, context_dim)

to_qkv_c instance-attribute

to_qkv_c = nn.Linear(context_dim, 3 * inner_dim)

forward

forward(
    x: Tensor,
    c: Tensor,
    mask: Tensor | None = None,
    rope: Tensor | None = None,
    c_rope: Tensor | None = None,
    c_mask: Tensor | None = None,
) -> tuple[Tensor, Tensor]

Rotary

Bases: Module

Rotary position frequencies for one stream, evaluated in fp32.

The frequencies live in a buffer so that a released checkpoint's own copy of them loads with strict=True. They are recomputed whenever the module has been cast below fp32: rounding either the frequencies or the integer positions to bf16 wrecks the phase at long positions.

base instance-attribute

base = base

dim instance-attribute

dim = dim

forward

forward(seq_len: int) -> Tensor

Return frequencies [1, 1, seq_len, dim] for positions 0..seq_len-1.

SingleBlock

Bases: Module

DiT block over the concatenated text and audio sequence.

attn instance-attribute

attn = Attention(
    dim=dim,
    heads=heads,
    dim_head=dim_head,
    attn_mask_enabled=attn_mask_enabled,
)

attn_norm instance-attribute

attn_norm = AdaLayerNorm(dim)

ff instance-attribute

ff = FeedForward(dim=dim, mult=ff_mult)

ff_norm instance-attribute

ff_norm = nn.LayerNorm(
    dim, elementwise_affine=False, eps=1e-06
)

forward

forward(
    x: Tensor, t: Tensor, mask: Tensor | None, rope: Tensor
) -> Tensor

TimeEmbedding

Bases: Module

Sinusoidal timestep features followed by a two-layer MLP.

freq_embed_dim instance-attribute

freq_embed_dim = freq_embed_dim

time_mlp instance-attribute

time_mlp = nn.Sequential(
    nn.Linear(freq_embed_dim, dim),
    nn.SiLU(),
    nn.Linear(dim, dim),
)

forward

forward(timestep: Tensor, scale: float = 1000.0) -> Tensor

dit_state_dict

dit_state_dict(
    checkpoint: Mapping[str, Tensor]
    | Iterable[tuple[str, Tensor]],
    prefix: str = "transformer.",
) -> dict[str, Tensor]

Select the DiT tensors out of a full AuK checkpoint and strip prefix.

Released checkpoints wrap the backbone under transformer. alongside the text-encoder layer-fusion parameters, which this module does not own.

sample_latents

sample_latents(
    dit: AuKTransformer,
    *,
    text: Tensor,
    c_mask: Tensor | None,
    ref: Tensor,
    ref_mask: Tensor | None,
    gen_frames: int,
    nfe: int = 32,
    cfg_strength: float = 1.0,
    sway_sampling_coef: float | None = None,
    t_grid: list[float] | None = None,
    seed: int | None = None,
    generator: Generator | None = None,
    latent_dim: int | None = None,
    device: device | str | None = None,
    dtype: dtype | None = None,
) -> Tensor

Integrate the flow from noise to audio latents with explicit Euler steps.

Single-request only: the batch dimension carries the CFG branches, not separate prompts, so no target padding mask is needed.

Parameters:

Name Type Description Default
dit AuKTransformer

The velocity model.

required
text Tensor

Pre-encoded text hidden states [1, nt, text_hidden_dim].

required
c_mask Tensor | None

Text padding mask [1, nt].

required
ref Tensor

Reference prompt latents [1, np, latent_dim].

required
ref_mask Tensor | None

Reference padding mask [1, np].

required
gen_frames int

Target latent frames to generate.

required
nfe int

Euler steps, ignored when t_grid is given.

32
cfg_strength float

Classifier-free guidance weight. Below 1e-5 the uncond branch is skipped entirely.

1.0
sway_sampling_coef float | None

Reshapes the uniform time grid towards t=0.

None
t_grid list[float] | None

Explicit timesteps, overriding nfe and the sway reshape.

None
seed int | None

Seeds the global RNG before drawing the noise. Ignored when generator is given.

None
generator Generator | None

Draws the initial noise from its own RNG, leaving the global one untouched.

None
latent_dim int | None

Latent channels; defaults to the model's.

None
device device | str | None

Device for the noise; defaults to the reference's.

None
dtype dtype | None

Dtype for the noise and the Euler accumulator; defaults to the reference's, which is the reference implementation's rule. Pass torch.float32 explicitly when serving half-precision weights. Drawing the noise at bf16 instead of fp32 moves a 32-step CFG trajectory by 0.025 relative MSE, roughly 17 times what bf16 weights themselves cost, because the draw consumes the generator differently and starts the ODE somewhere else. Keeping it fp32 costs a few tens of KB and no measurable time.

None

Returns:

Type Description
Tensor

The target latents at t=1, [1, gen_frames, latent_dim].