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. | required |
prefix | str | Unused; kept for the pipeline construction contract. | '' |
flash_t_grid instance-attribute ¶
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] |
|
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 | 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 |
|
project_text ¶
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,
)
)
decode ¶
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 ¶
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.
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 | 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 |