Skip to content

vllm_omni.diffusion.models.auk

AuK diffusion-stage components: the flow transformer, the codec, the pipeline.

Modules:

Name Description
alias_free_triton

Triton kernels for the AuK codec's alias-free SnakeBeta activation.

auk_transformer

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

auk_vae

BigVGAN-flow VAE for AuK: 24 kHz mono waveform <-> 50 Hz, 64-channel latents.

cudagraph_wrapper

Per-shape CUDA graph for one AuK DiT denoise step.

pipeline_auk

AuK audio-editing pipeline for vLLM-Omni.

vae_cudagraph

CUDA graph replay for the AuK codec decode.

AuKPipeline

Bases: Module, SupportAudioInput, SupportAudioOutput, SupportsComponentDiscovery

Instruction-driven audio generation and editing with AuK.

One request per forward: the rectified-flow ODE runs over the whole target span and the batch dimension carries the CFG branches, so there is nothing to share between requests yet.

Parameters:

Name Type Description Default
od_config OmniDiffusionConfig

OmniDiffusion configuration. od_config.model must be an assembled AuK directory (config.json, auk.safetensors, vae.safetensors), as produced by tools/prepare_auk_checkpoint.py.

required
prefix str

Unused; kept for the pipeline construction contract.

''

audio_sample_rate class-attribute

audio_sample_rate: int = 24000

cudagraph_wrapper instance-attribute

cudagraph_wrapper = AuKCUDAGraphWrapper(
    self.dit,
    enabled=not od_config.enforce_eager,
    max_graphs=max_dit_graphs,
)

default_cfg instance-attribute

default_cfg = float(defaults.get('cfg', 2.0))

default_nfe instance-attribute

default_nfe = int(defaults.get('nfe', 32))

default_sway instance-attribute

default_sway = defaults.get('sway', -1.0)

device instance-attribute

device = get_local_device()

dit instance-attribute

dit = self.dit.to(device=self.device).eval()

dtype instance-attribute

dtype = getattr(od_config, 'dtype', None) or torch.bfloat16

dummy_run_num_frames class-attribute

dummy_run_num_frames: int = 0

flash_t_grid instance-attribute

flash_t_grid = [
    float(t) for t in config.get("flash_t_grid") or ()
]

hop_size instance-attribute

hop_size = int(vae_config['downsample_rate'])

is_flash instance-attribute

is_flash = self.variant == 'flash'

latent_dim instance-attribute

latent_dim = int(vae_config['latent_dim'])

od_config instance-attribute

od_config = od_config

sample_rate instance-attribute

sample_rate = int(vae_config['target_sample_rate'])

support_audio_input class-attribute

support_audio_input: bool = True

support_audio_output class-attribute

support_audio_output: bool = True

supports_request_batch class-attribute instance-attribute

supports_request_batch = False

text_hidden_dim instance-attribute

text_hidden_dim = int(config['dit']['text_hidden_dim'])

vae instance-attribute

vae = self.vae.eval()

vae_decode instance-attribute

vae_decode = AuKVAEDecodeGraph(
    self.vae,
    enabled=not od_config.enforce_eager,
    **vae_decode_kwargs,
)

variant instance-attribute

variant = str(config.get('variant', 'base'))

weights_sources class-attribute

weights_sources: tuple = ()

forward

Generate one waveform.

Parameters:

Name Type Description Default
req DiffusionRequestBatch

Request batch holding exactly one request. The prompt carries prompt_embeds (the fused text condition [nt, 2048]), an optional multi_modal_data["audio"] source clip, and additional_information["auk"] with gen_seconds, sway, t_grid and vae_sample. num_inference_steps, guidance_scale and seed/generator come from the sampling params.

required

Returns:

Type Description
list[DiffusionOutput]

One DiffusionOutput whose output is a float32 mono waveform

list[DiffusionOutput]

[T] at 24 kHz, or the normalized target latents

list[DiffusionOutput]

[1, gen_frames, latent_dim] when output_type is latent.

load_weights

load_weights(
    weights: Iterable[tuple[str, Tensor]],
) -> set[str]

setup_compile

setup_compile() -> None

Compile the DiT as configured, then capture the codec decode buckets.

Defining this hook replaces the runner's generic transformer compile, so the DiT's diffusion_compile_granularity is honoured here: full compiles the whole denoise step; regional compiles the double- and single-stream blocks. Compilation is lazy and happens in the DiT CUDA graph's warm-up, before capture, so each graph replays the compiled kernels. The startup cost worth paying is the VAE decode, whose buckets are compiled and captured before the first request.

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.

AuKVAE

Bases: Module

AuK's audio codec: :meth:encode waveform to normalized latents, :meth:decode back.

Build from the released model.vae.model_init_kwargs with :meth:from_config, then :meth:load_weights. Latents are normalized by the checkpoint's global statistics, which is the space the rectified-flow transformer works in.

activation_post instance-attribute

activation_post = AliasFreeActivation(
    SnakeBeta(channels, alpha_logscale=snake_logscale),
    causal=act_causal,
)

audio_encoder instance-attribute

audio_encoder = Encoder(
    latent_dim=latent_dim,
    channels=tuple(downsample_channels),
    down_sample_factors=tuple(downsample_rates),
)

conv_post instance-attribute

conv_post = _weight_norm(
    Conv(channels, 1, 7, 1, causal=causal, bias=False)
)

conv_pre instance-attribute

conv_pre = _weight_norm(
    Conv(
        latent_dim,
        upsample_initial_channel,
        7,
        1,
        causal=False,
    )
)

hop_size instance-attribute

hop_size = math.prod(downsample_rates)

latent_dim instance-attribute

latent_dim = latent_dim

num_kernels instance-attribute

num_kernels = len(resblock_kernel_sizes)

num_upsamples instance-attribute

num_upsamples = len(upsample_rates)

resblocks instance-attribute

resblocks = nn.ModuleList()

sample_rate instance-attribute

sample_rate = sample_rate

ups instance-attribute

ups = nn.ModuleList()

decode

decode(latents: Tensor) -> Tensor

Decode normalized [B, Np, latent_dim] latents into a [B, Np * hop_size] waveform in [-1, 1].

decode_context_frames

decode_context_frames() -> tuple[int, int]

Latent frames of (left, right) context a decoded frame can depend on.

Summed layer by layer from the modules' own padding and kernel geometry, so it is an upper bound that holds for any weights: a window decoded on its own reproduces the full decode exactly except within these many frames of a window edge that is not a true clip edge. The decoder is causal apart from conv_pre and the band-limited upsamplers, so the right context is a few frames while the left one is dominated by the dilated convolutions of the first (lowest-rate) stage.

encode

encode(
    wav: Tensor,
    *,
    sample: bool = False,
    generator: Generator | None = None,
) -> Tensor

Encode a mono [1, T] waveform at :attr:sample_rate into [1, T / hop_size, latent_dim].

sample=False takes the posterior mean, which is what a reproducible pipeline wants; sample=True draws mean + randn * exp(log_std) from generator. Latents come back normalized by the checkpoint's global statistics.

from_config classmethod

from_config(cfg: dict) -> AuKVAE

Build from the released model_init_kwargs mapping, ignoring training-only entries.

load_weights

load_weights(
    path: str | Path,
    *,
    device: str | device | None = None,
    fold_weight_norm: bool = True,
) -> tuple[list[str], list[str]]

Load a released vae.safetensors, then fold weight norm away. Returns (missing, unexpected).

device defaults to where the module already lives, which is the placement that matters: the fold reduces weight_v to a norm, and a CPU reduction lands up to one ULP away from the accelerator's. The encoder amplifies its activations roughly fiftyfold, so that one ULP grows into a ~4e-4 relative drift by the latent head. Move the module first, then load.

remove_weight_norm

remove_weight_norm() -> None

Fold weight_g/weight_v into plain weights, on whichever device the module is on.

Idempotent, and meant to run once the checkpoint is in. See :meth:load_weights for why the device this runs on is a numeric decision rather than a convenience.

set_decode_fast_paths

set_decode_fast_paths(*, cached_filters: bool) -> None

Toggle the FIR modules' per-channel filter caching.

On by default. The switch exists so the cache's cost can be measured against the plain expand-per-call path. Flip it before any CUDA graph is captured: it reallocates the cached taps, and a captured graph keeps reading the old buffers.

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.

get_auk_post_process_func

get_auk_post_process_func(od_config: OmniDiffusionConfig)

Create the post-processing function for AuK audio output.

Tensor output types pass through; anything else becomes a numpy waveform. The sample rate is not attached here: the output formatter reads AuKPipeline.audio_sample_rate for audio-output pipelines.

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].