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