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. | required |
prefix | str | Unused; kept for the pipeline construction contract. | '' |
cudagraph_wrapper instance-attribute ¶
cudagraph_wrapper = AuKCUDAGraphWrapper(
self.dit,
enabled=not od_config.enforce_eager,
max_graphs=max_dit_graphs,
)
flash_t_grid instance-attribute ¶
vae_decode instance-attribute ¶
vae_decode = AuKVAEDecodeGraph(
self.vae,
enabled=not od_config.enforce_eager,
**vae_decode_kwargs,
)
forward ¶
forward(
req: DiffusionRequestBatch,
) -> list[DiffusionOutput]
Generate one waveform.
Parameters:
| Name | Type | Description | Default |
|---|---|---|---|
req | DiffusionRequestBatch | Request batch holding exactly one request. The prompt carries | required |
Returns:
| Type | Description |
|---|---|
list[DiffusionOutput] | One |
list[DiffusionOutput] |
|
list[DiffusionOutput] |
|
setup_compile ¶
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 | True |
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)
]
)
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)
]
)
clear_cache ¶
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 | required |
text | Tensor | Pre-encoded text hidden states | required |
time | Tensor | Flow-matching timestep, a scalar or | required |
mask | Tensor | None | Target padding mask | None |
c_mask | Tensor | None | Text padding mask | None |
ref | Tensor | None | Reference prompt latents | None |
ref_mask | Tensor | None | Reference padding mask | 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 | False |
cache | bool | Reuse the projected text across calls, for an ODE loop over fixed conditioning. Only the | False |
Returns:
| Type | Description |
|---|---|
Tensor | Velocity |
Tensor |
|
modulation_table ¶
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 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,
)
)
decode ¶
Decode normalized [B, Np, latent_dim] latents into a [B, Np * hop_size] waveform in [-1, 1].
decode_context_frames ¶
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 ¶
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 ¶
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 | required |
c_mask | Tensor | None | Text padding mask | required |
ref | Tensor | Reference prompt latents | required |
ref_mask | Tensor | None | Reference padding mask | required |
gen_frames | int | Target latent frames to generate. | required |
nfe | int | Euler steps, ignored when | 32 |
cfg_strength | float | Classifier-free guidance weight. Below | 1.0 |
sway_sampling_coef | float | None | Reshapes the uniform time grid towards | None |
t_grid | list[float] | None | Explicit timesteps, overriding | None |
seed | int | None | Seeds the global RNG before drawing the noise. Ignored when | 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 | None |
Returns:
| Type | Description |
|---|---|
Tensor | The target latents at |