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.
AliasFreeActivation ¶
Bases: Module
Oversample, apply the pointwise activation, decimate, so snake harmonics do not fold back.
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
]
)
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.
Conv ¶
ConvTranspose ¶
Downsample ¶
Encoder ¶
LowPass ¶
NormConv ¶
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)
]
)
SnakeBeta ¶
Bases: Module
x + sin^2(alpha * x) / beta with per-channel learned alpha and beta (log-scale in AuK).
Upsample ¶
Bases: Module
Band-limited interpolation by ratio. Always non-causal, matching how AuK builds it.