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.

RopeCache module-attribute

RopeCache = tuple[torch.Tensor, torch.Tensor]

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 | None, mod: Tensor | None = None
) -> tuple[Tensor, Tensor, Tensor, Tensor, Tensor]

mod is the precomputed linear(silu(emb)); emb is then unused.

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 | None, mod: Tensor | None = None
) -> Tensor

mod is the precomputed linear(silu(emb)); emb is then unused.

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: RopeCache | None = None,
    bias: Tensor | None = None,
) -> Tensor

bias is the precomputed key-padding bias for mask; built here when absent.

AuKStepContext dataclass

Per-request state of the denoise step, built by :meth:AuKTransformer.prepare.

Batch rows are branches copies of the request (cond first under CFG). target_mask covers the target frames only and has the request's batch size; the other masks cover their full sequences and every branch.

audio_mask instance-attribute

audio_mask: Tensor | None

branches instance-attribute

branches: int

c instance-attribute

c: Tensor

c_mask instance-attribute

c_mask: Tensor | None

joint_bias instance-attribute

joint_bias: Tensor | None

modulation class-attribute instance-attribute

modulation: Tensor | None = None

prompt instance-attribute

prompt: Tensor | None

rope_audio instance-attribute

rope_audio: RopeCache

rope_single instance-attribute

rope_single: RopeCache

rope_text instance-attribute

rope_text: RopeCache

single_bias instance-attribute

single_bias: Tensor | None

single_mask instance-attribute

single_mask: Tensor | None

target_mask instance-attribute

target_mask: Tensor | None

clone

clone() -> AuKStepContext

A copy whose tensors own fresh storage.

copy_

copy_(other: AuKStepContext) -> None

Overwrite these tensors in place with other's, for CUDA graph static inputs.

tensors

tensors() -> list[Tensor | None]

Every tensor field in a fixed order, rope pairs flattened.

AuKTransformer

Bases: Module

Flow-matching velocity predictor for AuK audio latents.

The forward pass is :meth:prepare (per request) followed by :meth:step (per timestep); an Euler loop calls the two separately.

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

attn_mask_enabled instance-attribute

attn_mask_enabled = attn_mask_enabled

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.

modulation_table

modulation_table(timesteps: Tensor) -> Tensor

adaLN modulations for every timestep, [len(timesteps), sum of widths].

The modulations depend only on the timestep, and a request's time grid is known before its first step. Computing them for the whole grid at once reads each modulation weight once per request instead of once per step, which at batch 1 is what those matrix-vector products cost.

prepare

prepare(
    text: Tensor,
    *,
    target_len: int,
    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,
    timesteps: Tensor | None = None,
) -> AuKStepContext

Everything a denoise step needs that does not depend on x or the time.

An Euler loop over fixed conditioning calls this once and then :meth:step per timestep, so the text projection, the reference prompt embedding, the padding masks and biases and the rotary tables are not recomputed on every step. Arguments are those of :meth:forward; target_len is the target frame count n. timesteps are the times the steps will run at (the grid without its end point); when given, their adaLN modulations are precomputed and :meth:step selects them by step_index.

project_text

project_text(text: Tensor) -> Tensor

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

step

step(
    x: Tensor,
    time: Tensor,
    ctx: AuKStepContext,
    step_index: Tensor | int | None = None,
) -> Tensor

Predict the velocity of x at time under a :meth:prepare context.

The target embedding does not depend on the CFG branch, so it is computed once and shared by both branches. When the context carries a modulation table, step_index selects this step's row of it and time is unused; the row broadcasts over the CFG branches.

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: RopeCache,
    c_rope: RopeCache,
    c_mask: Tensor | None,
    bias: Tensor | None = None,
    mod_c: Tensor | None = None,
    mod_x: Tensor | None = 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: RopeCache | None = None,
    c_rope: RopeCache | None = None,
    c_mask: Tensor | None = None,
    bias: Tensor | None = None,
) -> tuple[Tensor, Tensor]

bias is the precomputed :func:joint_key_mask bias; built here when absent.

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, mask: Tensor | None = None) -> Tensor

Return rotary frequencies, compressing positions across padding.

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: RopeCache,
    bias: Tensor | None = None,
    mod: Tensor | None = None,
) -> 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

build_time_grid

build_time_grid(
    *,
    nfe: int,
    sway_sampling_coef: float | None,
    t_grid: list[float] | None,
    device: device | str,
) -> Tensor

Build and validate the Euler schedule before a graph replay.

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.

joint_key_mask

joint_key_mask(
    mask: Tensor | None,
    c_mask: Tensor | None,
    text_len: int,
) -> Tensor | None

Key mask of joint attention over [audio | text]; None when the audio has no mask.

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,
    sampler: Callable[..., Tensor] | 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].