Skip to content

vllm_omni.diffusion.models.auk.pipeline_auk

AuK audio-editing pipeline for vLLM-Omni.

Turns a layer-fused text condition (produced by the encoder stage) plus an optional source clip into a 24 kHz waveform: the source clip is encoded to VAE latents, a rectified-flow DiT integrates the target latents conditioned on both, and the BigVGAN-flow decoder renders the waveform.

MAX_LATENT_FRAMES module-attribute

MAX_LATENT_FRAMES = 65536

logger module-attribute

logger = init_logger(__name__)

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]

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.