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_transformersandtorchdiffeqare 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_attnbackend 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_enableddefaults toTrue, the value the released config sets, rather than the reference'sFalse.
AdaLayerNorm ¶
Bases: Module
Timestep-conditioned modulation for a block: six chunks from one projection.
AdaLayerNormFinal ¶
Bases: Module
Timestep-conditioned modulation before the output projection.
forward ¶
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.
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.
copy_ ¶
copy_(other: AuKStepContext) -> None
Overwrite these tensors in place with other's, for CUDA graph static inputs.
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.
AudioEmbedding ¶
Bases: Module
Project audio latents to model width and add conv position information.
ConvPosEmbedding ¶
Bases: Module
Depthwise-grouped conv stack that adds local position information.
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,
)
FeedForward ¶
JointAttention ¶
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.
SingleBlock ¶
TimeEmbedding ¶
Bases: Module
Sinusoidal timestep features followed by a two-layer MLP.
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 | 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 |