Skip to content

vllm_omni.diffusion.models.auk

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

Modules:

Name Description
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.

pipeline_auk

AuK audio-editing pipeline for vLLM-Omni.

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

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

vae instance-attribute

vae = self.vae.eval()

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]

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.

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

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.

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