Skip to content

vllm_omni.diffusion.models.auk.auk_vae

BigVGAN-flow VAE for AuK: 24 kHz mono waveform <-> 50 Hz, 64-channel latents.

Inference only. The reference module also carries a residual coupling flow, a KL term and a discriminator; all three are training-only, so the released flow.* tensors are skipped (see :meth:AuKVAE.load_weights).

Parameter and buffer names follow the reference release, so the published vae.safetensors loads without a key map. That is also why the convolutions are built with the legacy weight_norm parametrization (weight_g / weight_v): :meth:remove_weight_norm folds them into plain weights once the checkpoint is in.

TRAINING_ONLY_PREFIX module-attribute

TRAINING_ONLY_PREFIX = 'flow.'

logger module-attribute

logger = init_logger(__name__)

AliasFreeActivation

Bases: Module

Oversample, apply the pointwise activation, decimate, so snake harmonics do not fold back.

act instance-attribute

act = activation

downsample instance-attribute

downsample = Downsample(
    down_ratio, down_kernel_size, causal=causal
)

upsample instance-attribute

upsample = Upsample(up_ratio, up_kernel_size)

forward

forward(x: Tensor) -> Tensor

AmpBlock

Bases: Module

Anti-aliased multi-periodicity residual block: alternating alias-free snake and dilated convs.

activations instance-attribute

activations = nn.ModuleList(
    [
        AliasFreeActivation(
            SnakeBeta(
                channels, alpha_logscale=snake_logscale
            ),
            causal=act_causal,
        )
        for _ in range(
            len(self.convs1) + len(self.convs2)
        )
    ]
)

convs1 instance-attribute

convs1 = nn.ModuleList(
    [
        _weight_norm(
            Conv(
                channels,
                channels,
                kernel_size,
                1,
                dilation=d,
                causal=causal,
            )
        )
        for d in dilations
    ]
)

convs2 instance-attribute

convs2 = nn.ModuleList(
    [
        _weight_norm(
            Conv(
                channels,
                channels,
                kernel_size,
                1,
                dilation=1,
                causal=causal,
            )
        )
        for _ in dilations
    ]
)

forward

forward(x: Tensor) -> Tensor

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,
    )
)

hop_size instance-attribute

hop_size = math.prod(downsample_rates)

latent_dim instance-attribute

latent_dim = latent_dim

num_kernels instance-attribute

num_kernels = len(resblock_kernel_sizes)

num_upsamples instance-attribute

num_upsamples = len(upsample_rates)

resblocks instance-attribute

resblocks = nn.ModuleList()

sample_rate instance-attribute

sample_rate = sample_rate

ups instance-attribute

ups = nn.ModuleList()

decode

decode(latents: Tensor) -> Tensor

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

from_config(cfg: dict) -> AuKVAE

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

remove_weight_norm() -> None

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.

Conv

Bases: Conv1d

1-D convolution over [B, C, T]; causal swaps centred padding for a left pad.

causal instance-attribute

causal = causal

left_padding instance-attribute

left_padding = dilation * (kernel_size - 1) if causal else 0

forward

forward(x: Tensor) -> Tensor

ConvTranspose

Bases: ConvTranspose1d

Transposed 1-D convolution; causal drops the trailing stride samples instead of padding.

causal instance-attribute

causal = causal

trim instance-attribute

trim = stride

forward

forward(x: Tensor) -> Tensor

Downsample

Bases: Module

Decimation by ratio behind a low-pass; lowpass holds the filter buffer.

lowpass instance-attribute

lowpass = LowPass(
    cutoff=0.5 / ratio,
    half_width=0.6 / ratio,
    stride=ratio,
    kernel_size=kernel_size,
    causal=causal,
)

forward

forward(x: Tensor) -> Tensor

Encoder

Bases: Module

Strided conv encoder over [B, 1, T], emitting [B, 2 * latent_dim, T / hop] mean/log-std pairs.

generator instance-attribute

generator = nn.Sequential(*layers)

forward

forward(x: Tensor) -> Tensor

LowPass

Bases: Module

Fixed FIR low-pass with replicate padding, applied per channel.

pad_left instance-attribute

pad_left = kernel_size // 2 - int(even)

pad_right instance-attribute

pad_right = kernel_size // 2

stride instance-attribute

stride = stride

forward

forward(x: Tensor) -> Tensor

NormConv

Bases: Module

Weight-normalized conv held under layer; the attribute name is load-bearing for the checkpoint.

layer instance-attribute

layer = _weight_norm(
    nn.Conv1d(
        in_channels,
        out_channels,
        kernel_size,
        stride=stride,
        padding=(kernel_size - 1) // 2,
    )
)

forward

forward(x: Tensor) -> Tensor

ResStack

Bases: Module

Stack of dilated residual conv pairs. The leaky slope is torch's 0.01 default, as upstream.

layers instance-attribute

layers = nn.ModuleList(
    [
        nn.Sequential(
            nn.LeakyReLU(),
            _weight_norm(
                nn.Conv1d(
                    channels,
                    channels,
                    kernel_size,
                    dilation=base**i,
                    padding=base**i,
                )
            ),
            nn.LeakyReLU(),
            _weight_norm(
                nn.Conv1d(
                    channels,
                    channels,
                    kernel_size,
                    dilation=1,
                    padding=1,
                )
            ),
        )
        for i in range(nums)
    ]
)

forward

forward(x: Tensor) -> Tensor

SnakeBeta

Bases: Module

x + sin^2(alpha * x) / beta with per-channel learned alpha and beta (log-scale in AuK).

alpha instance-attribute

alpha = nn.Parameter(init.clone())

alpha_logscale instance-attribute

alpha_logscale = alpha_logscale

beta instance-attribute

beta = nn.Parameter(init.clone())

eps instance-attribute

eps = 1e-09

forward

forward(x: Tensor) -> Tensor

Upsample

Bases: Module

Band-limited interpolation by ratio. Always non-causal, matching how AuK builds it.

pad instance-attribute

pad = kernel_size // ratio - 1

pad_left instance-attribute

pad_left = self.pad * ratio + (kernel_size - ratio) // 2

pad_right instance-attribute

pad_right = (
    self.pad * ratio + (kernel_size - ratio + 1) // 2
)

ratio instance-attribute

ratio = ratio

forward

forward(x: Tensor) -> Tensor