Skip to content

vllm_omni.model_executor.models.audio8_tts.codec

Audio8 TTS neural audio codec (44.1 kHz, 10 codebooks).

Vendored from the reference modeling_arktts_codec.py shipped with Audio8/Audio8-TTS-Preview-0.6b (Apache-2.0). The module tree and attribute names are reproduced verbatim so codec.pth loads under strict=True.

Inference-only deviations: no torch.jit.script on Snake (it fights torch.compile / CUDA graphs), RoPE cached in a plain dict instead of a non-persistent buffer (the pattern that leaves the reference LM with uninitialised RoPE under from_pretrained on transformers >= 5), and :func:build_arktts_codec prunes the unused encoder or decoder half per stage.

ArkttsCausalConv1d

Bases: Module

conv instance-attribute

conv = nn.Conv1d(
    in_channels,
    out_channels,
    kernel_size,
    stride=stride,
    dilation=dilation,
    groups=groups,
)

kernel_size instance-attribute

kernel_size = (kernel_size - 1) * dilation + 1

padding instance-attribute

padding = self.kernel_size - self.stride

stride instance-attribute

stride = stride

apply_weight_norm

apply_weight_norm() -> ArkttsCausalConv1d

forward

forward(x: Tensor) -> Tensor

ArkttsCausalConvTranspose1d

Bases: Module

conv instance-attribute

conv = nn.ConvTranspose1d(
    in_channels,
    out_channels,
    kernel_size,
    stride=stride,
    dilation=dilation,
)

kernel_size instance-attribute

kernel_size = kernel_size

stride instance-attribute

stride = stride

apply_weight_norm

apply_weight_norm() -> ArkttsCausalConvTranspose1d

forward

forward(x: Tensor) -> Tensor

ArkttsCodec

Bases: Module

Audio8 TTS codec: waveform <-> 10 codebooks at ~21.5 frames/s.

decoder instance-attribute

decoder = ArkttsDecoder()

encoder instance-attribute

encoder = ArkttsEncoder()

frame_length class-attribute instance-attribute

frame_length = 2048

quantizer instance-attribute

quantizer = ArkttsDownsampleQuantizer(
    post_n_layer=post_n_layer,
    post_n_head=post_n_head,
    post_n_local_heads=post_n_local_heads,
    post_intermediate_size=post_intermediate_size,
)

sample_rate class-attribute instance-attribute

sample_rate = 44100

decode

decode(codes: Tensor) -> Tensor

Decode [B, 10, frames] to [B, 1, samples].

encode

encode(
    audio: Tensor, audio_lengths: Tensor | None = None
) -> tuple[Tensor, Tensor]

Encode [B, 1, samples] (or [B, samples]) to [B, 10, frames].

ArkttsCodecAttention

Bases: Module

head_dim instance-attribute

head_dim = config.head_dim

n_head instance-attribute

n_head = config.n_head

n_local_heads instance-attribute

n_local_heads = config.n_local_heads

wo instance-attribute

wo = nn.Linear(
    config.head_dim * config.n_head, config.dim, bias=False
)

wqkv instance-attribute

wqkv = nn.Linear(config.dim, total, bias=False)

forward

forward(
    x: Tensor, rope_values: Tensor, mask: Tensor
) -> Tensor

ArkttsCodecFeedForward

Bases: Module

w1 instance-attribute

w1 = nn.Linear(
    config.dim, config.intermediate_size, bias=False
)

w2 instance-attribute

w2 = nn.Linear(
    config.intermediate_size, config.dim, bias=False
)

w3 instance-attribute

w3 = nn.Linear(
    config.dim, config.intermediate_size, bias=False
)

forward

forward(x: Tensor) -> Tensor

ArkttsCodecLayerScale

Bases: Module

gamma instance-attribute

gamma = nn.Parameter(init_values * torch.ones(dim))

forward

forward(x: Tensor) -> Tensor

ArkttsCodecRMSNorm

Bases: Module

eps instance-attribute

eps = float(eps)

weight instance-attribute

weight = nn.Parameter(torch.ones(dim))

forward

forward(x: Tensor) -> Tensor

ArkttsCodecTransformerBlock

Bases: Module

attention instance-attribute

attention = ArkttsCodecAttention(config)

attention_layer_scale instance-attribute

attention_layer_scale = ArkttsCodecLayerScale(config.dim)

attention_norm instance-attribute

attention_norm = ArkttsCodecRMSNorm(
    config.dim, config.norm_eps
)

feed_forward instance-attribute

feed_forward = ArkttsCodecFeedForward(config)

ffn_layer_scale instance-attribute

ffn_layer_scale = ArkttsCodecLayerScale(config.dim)

ffn_norm instance-attribute

ffn_norm = ArkttsCodecRMSNorm(config.dim, config.norm_eps)

forward

forward(
    x: Tensor, rope_values: Tensor, mask: Tensor
) -> Tensor

ArkttsCodecTransformerConfig

Plain option holder for the codec's transformer blocks.

Deliberately not a dataclass / msgspec Struct: it mirrors a reference dataclass whose only job is to carry constructor arguments, and it is instantiated from Python literals inside this module only.

channels_first instance-attribute

channels_first = bool(channels_first)

dim instance-attribute

dim = int(dim)

head_dim instance-attribute

head_dim = int(head_dim)

intermediate_size instance-attribute

intermediate_size = int(intermediate_size)

n_head instance-attribute

n_head = int(n_head)

n_layer instance-attribute

n_layer = int(n_layer)

n_local_heads instance-attribute

n_local_heads = (
    int(n_head)
    if n_local_heads == -1
    else int(n_local_heads)
)

norm_eps instance-attribute

norm_eps = float(norm_eps)

rope_base instance-attribute

rope_base = float(rope_base)

ArkttsCodecWindowTransformer

Bases: Module

Causal transformer with a sliding attention window.

channels_first instance-attribute

channels_first = config.channels_first

head_dim instance-attribute

head_dim = config.head_dim

input_proj instance-attribute

input_proj = (
    nn.Linear(input_dim, config.dim)
    if input_dim != config.dim
    else nn.Identity()
)

layers instance-attribute

layers = nn.ModuleList(
    [
        ArkttsCodecTransformerBlock(config)
        for _ in range(config.n_layer)
    ]
)

look_ahead_conv instance-attribute

look_ahead_conv = nn.Identity()

norm instance-attribute

norm = ArkttsCodecRMSNorm(config.dim, config.norm_eps)

output_proj instance-attribute

output_proj = (
    nn.Linear(config.dim, input_dim)
    if input_dim != config.dim
    else nn.Identity()
)

rope_base instance-attribute

rope_base = config.rope_base

window_size instance-attribute

window_size = window_size

forward

forward(x: Tensor) -> Tensor

ArkttsConvNeXtBlock

Bases: Module

act instance-attribute

act = nn.GELU()

dwconv instance-attribute

dwconv = ArkttsCausalConv1d(
    dim, dim, kernel_size=7, groups=dim
)

gamma instance-attribute

gamma = nn.Parameter(1e-06 * torch.ones(dim))

norm instance-attribute

norm = nn.LayerNorm(dim, eps=1e-06)

pwconv1 instance-attribute

pwconv1 = nn.Linear(dim, 4 * dim)

pwconv2 instance-attribute

pwconv2 = nn.Linear(4 * dim, dim)

forward

forward(x: Tensor) -> Tensor

ArkttsDecoder

Bases: Module

model instance-attribute

model = nn.Sequential(*modules)

forward

forward(x: Tensor) -> Tensor

ArkttsDecoderBlock

Bases: Module

block instance-attribute

block = nn.Sequential(
    ArkttsSnake1d(input_dim),
    _causal_wn_transpose(
        input_dim,
        output_dim,
        kernel_size=2 * stride,
        stride=stride,
    ),
    ArkttsResidualUnit(output_dim, 1),
    ArkttsResidualUnit(output_dim, 3),
    ArkttsResidualUnit(output_dim, 9),
)

forward

forward(x: Tensor) -> Tensor

ArkttsDownsampleQuantizer

Bases: Module

1 semantic codebook (4096) + 9 residual codebooks (1024), 4x downsampled.

downsample instance-attribute

downsample = nn.Sequential(
    nn.Sequential(
        ArkttsCausalConv1d(
            1024, 1024, kernel_size=2, stride=2
        ),
        ArkttsConvNeXtBlock(1024),
    ),
    nn.Sequential(
        ArkttsCausalConv1d(
            1024, 1024, kernel_size=2, stride=2
        ),
        ArkttsConvNeXtBlock(1024),
    ),
)

post_module instance-attribute

post_module = ArkttsCodecWindowTransformer(
    ArkttsCodecTransformerConfig(
        n_layer=post_n_layer,
        n_head=post_n_head,
        n_local_heads=post_n_local_heads,
        dim=1024,
        intermediate_size=post_intermediate_size,
    ),
    1024,
    window_size=128,
)

pre_module instance-attribute

pre_module = ArkttsCodecWindowTransformer(
    ArkttsCodecTransformerConfig(
        n_layer=8,
        n_head=16,
        dim=1024,
        intermediate_size=3072,
    ),
    1024,
    window_size=128,
)

quantizer instance-attribute

quantizer = ArkttsResidualQuantizer(1024, 9, 1024, 8)

semantic_predictor_module instance-attribute

semantic_predictor_module = nn.Identity()

semantic_quantizer instance-attribute

semantic_quantizer = ArkttsResidualQuantizer(
    1024, 1, 4096, 8
)

upsample instance-attribute

upsample = nn.Sequential(
    nn.Sequential(
        ArkttsCausalConvTranspose1d(
            1024, 1024, kernel_size=2, stride=2
        ),
        ArkttsConvNeXtBlock(1024),
    ),
    nn.Sequential(
        ArkttsCausalConvTranspose1d(
            1024, 1024, kernel_size=2, stride=2
        ),
        ArkttsConvNeXtBlock(1024),
    ),
)

decode

decode(indices: Tensor) -> Tensor

forward

forward(z: Tensor) -> tuple[Tensor, Tensor]

ArkttsEncoder

Bases: Module

block instance-attribute

block = nn.Sequential(*modules)

forward

forward(x: Tensor) -> Tensor

ArkttsEncoderBlock

Bases: Module

block instance-attribute

block = nn.Sequential(*modules)

forward

forward(x: Tensor) -> Tensor

ArkttsResidualQuantizer

Bases: Module

codebook_size instance-attribute

codebook_size = int(codebook_size)

n_codebooks instance-attribute

n_codebooks = int(n_codebooks)

quantizers instance-attribute

quantizers = nn.ModuleList(
    [
        ArkttsVectorQuantizer(
            input_dim, codebook_size, codebook_dim
        )
        for _ in range(n_codebooks)
    ]
)

forward

forward(z: Tensor) -> tuple[Tensor, Tensor]

from_codes

from_codes(codes: Tensor) -> Tensor

ArkttsResidualUnit

Bases: Module

block instance-attribute

block = nn.Sequential(
    ArkttsSnake1d(dim),
    _causal_wn_conv(
        dim, dim, kernel_size=7, dilation=dilation
    ),
    ArkttsSnake1d(dim),
    _causal_wn_conv(dim, dim, kernel_size=1),
)

forward

forward(x: Tensor) -> Tensor

ArkttsSnake1d

Bases: Module

alpha instance-attribute

alpha = nn.Parameter(torch.ones(1, channels, 1))

forward

forward(x: Tensor) -> Tensor

ArkttsVectorQuantizer

Bases: Module

codebook instance-attribute

codebook = nn.Embedding(codebook_size, codebook_dim)

codebook_dim instance-attribute

codebook_dim = int(codebook_dim)

codebook_size instance-attribute

codebook_size = int(codebook_size)

in_proj instance-attribute

in_proj = legacy_weight_norm(
    nn.Conv1d(input_dim, codebook_dim, kernel_size=1)
)

out_proj instance-attribute

out_proj = legacy_weight_norm(
    nn.Conv1d(codebook_dim, input_dim, kernel_size=1)
)

decode_code

decode_code(indices: Tensor) -> Tensor

decode_latents

decode_latents(latents: Tensor) -> tuple[Tensor, Tensor]

forward

forward(z: Tensor) -> tuple[Tensor, Tensor]

build_arktts_codec

build_arktts_codec(
    *,
    post_n_layer: int = 8,
    post_n_head: int = 16,
    post_n_local_heads: int = 8,
    post_intermediate_size: int = 1216,
) -> ArkttsCodec

Construct the codec with uninitialised weights, on CPU, in eval mode.