Skip to content

Models

RVQ autoencoder model architecture using heliaEDGE components.

compressionkit.models.rvq_autoencoder.build_rvq_autoencoder(frame_size, *, embedding_dim=16, latent_width=256, in_ch=1, out_ch=1, base_filters=32, multiplier=1.25, num_levels=2, beta=0.25, num_stages=4, encoder_block_norm='batch', encoder_head_norm='none', decoder_block_norm='none', decoder_head_norm='layer', use_residual=False, use_ema=False, ema_decay=0.99, encoder_type='default', expand_ratio=4.0, causal=False, discard_tail=0, bottleneck_type='rvq', fsq_levels=None, decoder_type='default', decoder_state_size=32, decoder_num_ssm_blocks=2, hier_detail_scale=0.25, revive_dead_codes=False, revive_threshold=0.03, kmeans_init=False, structured_dropout=False, dropout_levels=None, decoder_activation='relu', encoder_blocks_per_stage=1, codebook_sizes=None, prefix_loss_weights=None, prefix_loss_target='lowpass', prefix_loss_lowpass_kernel=9, prefix_loss_initial_scale=1.0)

Build encoder, bottleneck, decoder, and composite VQAutoencoder.

Returns:

Type Description
Model

(encoder, bottleneck, decoder, model) where model is a

ResidualVectorQuantizer | EmaResidualVectorQuantizer | FiniteScalarQuantizer

helia_edge.trainers.VQAutoencoder wrapping the three components.

Source code in compressionkit/models/rvq_autoencoder.py
def build_rvq_autoencoder(
    frame_size: int,
    *,
    embedding_dim: int = 16,
    latent_width: int = 256,
    in_ch: int = 1,
    out_ch: int = 1,
    base_filters: int = 32,
    multiplier: float = 1.25,
    num_levels: int = 2,
    beta: float = 0.25,
    num_stages: int = 4,
    encoder_block_norm: str = "batch",
    encoder_head_norm: str = "none",
    decoder_block_norm: str = "none",
    decoder_head_norm: str = "layer",
    use_residual: bool = False,
    use_ema: bool = False,
    ema_decay: float = 0.99,
    encoder_type: str = "default",
    expand_ratio: float = 4.0,
    causal: bool = False,
    discard_tail: int = 0,
    bottleneck_type: str = "rvq",
    fsq_levels: list[int] | None = None,
    decoder_type: str = "default",
    decoder_state_size: int = 32,
    decoder_num_ssm_blocks: int = 2,
    hier_detail_scale: float = 0.25,
    revive_dead_codes: bool = False,
    revive_threshold: float = 0.03,
    kmeans_init: bool = False,
    structured_dropout: bool = False,
    dropout_levels: list[int] | None = None,
    decoder_activation: str = "relu",
    encoder_blocks_per_stage: int = 1,
    codebook_sizes: list[int] | None = None,
    prefix_loss_weights: list[float] | None = None,
    prefix_loss_target: str = "lowpass",
    prefix_loss_lowpass_kernel: int = 9,
    prefix_loss_initial_scale: float = 1.0,
) -> tuple[
    keras.Model,
    ResidualVectorQuantizer | EmaResidualVectorQuantizer | FiniteScalarQuantizer,
    keras.Model,
    VQAutoencoder,
]:
    """Build encoder, bottleneck, decoder, and composite VQAutoencoder.

    Returns:
        ``(encoder, bottleneck, decoder, model)`` where *model* is a
        ``helia_edge.trainers.VQAutoencoder`` wrapping the three components.
    """
    downsample_factor = 2**num_stages
    if frame_size % downsample_factor != 0:
        raise ValueError(f"frame_size ({frame_size}) must be divisible by 2**num_stages ({downsample_factor})")

    bottleneck_type = bottleneck_type.lower()
    if bottleneck_type == "fsq":
        if not fsq_levels:
            raise ValueError("bottleneck_type='fsq' requires non-empty fsq_levels list")
        if embedding_dim != len(fsq_levels):
            embedding_dim = len(fsq_levels)
    elif bottleneck_type not in ("rvq", "significance"):
        raise ValueError(f"Unknown bottleneck_type: {bottleneck_type!r}")

    # --- Encoder ---
    if encoder_type == "soundstream":
        encoder = build_soundstream_encoder(
            input_len=frame_size,
            in_ch=in_ch,
            embedding_dim=embedding_dim,
            base_filters=base_filters,
            multiplier=multiplier,
            num_stages=num_stages,
            norm=encoder_block_norm,
            head_norm=encoder_head_norm,
        )
    elif encoder_type == "mlp":
        encoder = build_mlp_encoder(
            input_len=frame_size,
            in_ch=in_ch,
            embedding_dim=embedding_dim,
            num_stages=num_stages,
            hidden_dim=int(base_filters * multiplier),
            num_layers=encoder_blocks_per_stage,
        )
    elif encoder_type == "transformer":
        encoder = build_transformer_encoder(
            input_len=frame_size,
            in_ch=in_ch,
            embedding_dim=embedding_dim,
            num_stages=num_stages,
            d_model=int(base_filters * multiplier),
            num_heads=4,
            num_layers=encoder_blocks_per_stage,
            ff_dim=int(base_filters * multiplier * 2),
        )
    elif encoder_type == "inverted_residual":
        encoder = build_encoder_2d_invres(
            input_len=frame_size,
            in_ch=in_ch,
            base=base_filters,
            embedding_dim=embedding_dim,
            multiplier=multiplier,
            num_stages=num_stages,
            block_norm=encoder_block_norm,
            head_norm=encoder_head_norm,
            expand_ratio=expand_ratio,
            causal=causal,
            discard_tail=discard_tail,
        )
    else:
        encoder = build_encoder_2d(
            input_len=frame_size,
            in_ch=in_ch,
            base=base_filters,
            embedding_dim=embedding_dim,
            multiplier=multiplier,
            num_stages=num_stages,
            block_norm=encoder_block_norm,
            head_norm=encoder_head_norm,
            use_residual=use_residual,
            blocks_per_stage=encoder_blocks_per_stage,
        )

    # --- Decoder ---
    if decoder_type == "soundstream":
        decoder = build_soundstream_decoder(
            output_len=frame_size,
            out_ch=out_ch,
            embedding_dim=embedding_dim,
            base_filters=base_filters,
            multiplier=multiplier,
            num_stages=num_stages,
            norm=decoder_block_norm,
        )
    elif decoder_type == "mlp":
        decoder = build_mlp_decoder(
            output_len=frame_size,
            out_ch=out_ch,
            embedding_dim=embedding_dim,
            num_stages=num_stages,
            hidden_dim=int(base_filters * multiplier),
            num_layers=encoder_blocks_per_stage,
        )
    elif decoder_type == "transformer":
        decoder = build_transformer_decoder(
            output_len=frame_size,
            out_ch=out_ch,
            embedding_dim=embedding_dim,
            num_stages=num_stages,
            d_model=int(base_filters * multiplier),
            num_heads=4,
            num_layers=encoder_blocks_per_stage,
            ff_dim=int(base_filters * multiplier * 2),
        )
    elif decoder_type == "ssm":
        decoder = build_decoder_2d_ssm(
            output_len=frame_size,
            out_ch=out_ch,
            base=base_filters,
            embedding_dim=embedding_dim,
            multiplier=multiplier,
            num_stages=num_stages,
            head_norm=decoder_head_norm,
            state_size=decoder_state_size,
            num_ssm_blocks=decoder_num_ssm_blocks,
        )
    elif decoder_type in {"hierarchical", "hierarchical_hybrid"}:
        decoder = build_hierarchical_decoder_2d(
            output_len=frame_size,
            out_ch=out_ch,
            base=base_filters,
            embedding_dim=embedding_dim,
            num_levels=num_levels,
            multiplier=multiplier,
            num_stages=num_stages,
            decoder_block_norm=decoder_block_norm,
            head_norm=decoder_head_norm,
            use_residual=use_residual,
            activation=decoder_activation,
            detail_scale=hier_detail_scale,
            include_sum_input=(decoder_type == "hierarchical_hybrid"),
        )
    elif decoder_type == "hierarchical_adaptor":
        decoder = build_hierarchical_adaptor_decoder_2d(
            output_len=frame_size,
            out_ch=out_ch,
            base=base_filters,
            embedding_dim=embedding_dim,
            num_levels=num_levels,
            multiplier=multiplier,
            num_stages=num_stages,
            decoder_block_norm=decoder_block_norm,
            head_norm=decoder_head_norm,
            use_residual=use_residual,
            activation=decoder_activation,
            detail_scale=hier_detail_scale,
        )
    else:
        decoder = build_decoder_2d(
            output_len=frame_size,
            out_ch=out_ch,
            base=base_filters,
            embedding_dim=embedding_dim,
            multiplier=multiplier,
            num_stages=num_stages,
            decoder_block_norm=decoder_block_norm,
            head_norm=decoder_head_norm,
            use_residual=use_residual,
            activation=decoder_activation,
        )

    # --- Bottleneck ---
    if bottleneck_type == "fsq":
        bottleneck = FiniteScalarQuantizer(levels=list(fsq_levels))
        model = VQAutoencoder(
            encoder=encoder,
            vq=bottleneck,
            decoder=decoder,
            name=f"FSQAE_2D_ds{downsample_factor}",
        )
        return encoder, bottleneck, decoder, model

    if bottleneck_type == "significance":
        latent_len = frame_size // downsample_factor
        num_latents = latent_len * embedding_dim
        # keep_ratio: what fraction of latent values to keep.
        # Compute from bit budget: at input_bits/CR target, each kept value costs ~quant_bits
        # Default: keep half (can be tuned via beta repurposed as keep_ratio if >1, else rate_lambda)
        # We use num_levels to encode quant_bits, and beta as rate_lambda
        quant_bits = max(4, min(16, num_levels * 8 if num_levels >= 1 else 8))
        # Approximate keep_ratio for target CR:
        # budget_bits = frame_size * 16 / CR, where CR = frame_size / (latent_len * embedding_dim * keep_ratio * quant_bits / 16)
        # For simplicity: keep_ratio = budget_bits / (num_latents * quant_bits)
        # At 4× CR: budget = 320*16/4 = 1280, num_latents=160, qbits=8 → keep_ratio = 1280/(160*8) = 1.0
        # At 4× CR: budget = 320*16/4 = 1280, num_latents=320, qbits=8 → keep_ratio = 1280/(320*8) = 0.5
        keep_ratio = min(1.0, (frame_size * 16.0 / 4.0) / (num_latents * quant_bits))
        bottleneck = SignificanceQuantizer(
            keep_ratio=keep_ratio,
            quant_bits=quant_bits,
            rate_lambda=beta,
            temperature=0.1,
        )
        model = VQAutoencoder(
            encoder=encoder,
            vq=bottleneck,
            decoder=decoder,
            name=f"SigAE_2D_ds{downsample_factor}",
        )
        return encoder, bottleneck, decoder, model

    num_embeddings: int | list[int] = codebook_sizes if codebook_sizes else latent_width

    if use_ema:
        rvq = EmaResidualVectorQuantizer(
            num_levels=num_levels,
            num_embeddings=num_embeddings,
            embedding_dim=embedding_dim,
            beta=beta,
            ema_decay=ema_decay,
            revive_dead_codes=revive_dead_codes,
            revive_threshold=revive_threshold,
            kmeans_init=kmeans_init,
            structured_dropout=structured_dropout,
            dropout_levels=dropout_levels,
        )
    else:
        rvq = ResidualVectorQuantizer(
            num_levels=num_levels,
            num_embeddings=num_embeddings,
            embedding_dim=embedding_dim,
            beta=beta,
        )

    # --- Model class selection ---
    if decoder_type in {"hierarchical", "hierarchical_hybrid", "hierarchical_adaptor"}:
        if not use_ema:
            raise ValueError("hierarchical decoder types currently require use_ema=True")
        model_cls = HierarchicalRVQAutoencoder
    else:
        model_cls = PrefixSupervisedVQAutoencoder if prefix_loss_weights else VQAutoencoder
    model_kwargs = {
        "encoder": encoder,
        "vq": rvq,
        "decoder": decoder,
        "name": f"RVQAE_2D_ds{downsample_factor}",
    }
    if prefix_loss_weights:
        model_kwargs.update(
            {
                "prefix_loss_weights": prefix_loss_weights,
                "prefix_loss_target": prefix_loss_target,
                "prefix_loss_lowpass_kernel": prefix_loss_lowpass_kernel,
                "prefix_loss_initial_scale": prefix_loss_initial_scale,
            }
        )
    if decoder_type in {"hierarchical_hybrid", "hierarchical_adaptor"}:
        model_kwargs["include_summed_latent"] = True
    model = model_cls(**model_kwargs)

    return encoder, rvq, decoder, model

compressionkit.models.rvq_autoencoder.build_encoder_2d(input_len=2048, in_ch=1, base=32, embedding_dim=16, multiplier=1.25, num_stages=4, block_norm='batch', head_norm='none', use_residual=False, blocks_per_stage=1)

Build a configurable encoder with 2**num_stages downsampling.

Parameters:

Name Type Description Default
input_len int

Number of input time samples.

2048
in_ch int

Number of input channels.

1
base int

Base filter count for the first stage.

32
embedding_dim int

Latent channel dimension after projection.

16
multiplier float

Filter count multiplier per stage.

1.25
num_stages int

Number of stride-2 downsampling stages.

4
block_norm str

Normalization mode for conv blocks.

'batch'
head_norm str

Normalization mode for the final projection.

'none'
use_residual bool

If True, add shortcut connections to each stage.

False
blocks_per_stage int

Number of blocks per stage. First block does stride-2, additional blocks run at stride-1 for extra capacity.

1
Source code in compressionkit/models/encoder.py
def build_encoder_2d(
    input_len: int = 2048,
    in_ch: int = 1,
    base: int = 32,
    embedding_dim: int = 16,
    multiplier: float = 1.25,
    num_stages: int = 4,
    block_norm: str = "batch",
    head_norm: str = "none",
    use_residual: bool = False,
    blocks_per_stage: int = 1,
) -> keras.Model:
    """Build a configurable encoder with ``2**num_stages`` downsampling.

    Args:
        input_len: Number of input time samples.
        in_ch: Number of input channels.
        base: Base filter count for the first stage.
        embedding_dim: Latent channel dimension after projection.
        multiplier: Filter count multiplier per stage.
        num_stages: Number of stride-2 downsampling stages.
        block_norm: Normalization mode for conv blocks.
        head_norm: Normalization mode for the final projection.
        use_residual: If True, add shortcut connections to each stage.
        blocks_per_stage: Number of blocks per stage. First block does stride-2,
            additional blocks run at stride-1 for extra capacity.
    """
    downsample_factor = 2**num_stages
    if num_stages < 1:
        raise ValueError(f"num_stages must be >= 1, got {num_stages}")
    if input_len % downsample_factor != 0:
        raise ValueError(f"input_len ({input_len}) must be divisible by 2**num_stages ({downsample_factor})")

    conv_fn = res_conv2d_block if use_residual else conv2d_block
    dw_fn = res_depthwise2d_block if use_residual else depthwise2d_block

    inp = keras.layers.Input(shape=(1, input_len, in_ch), name="enc_in")
    x = inp
    filters = base
    conv_stages = min(2, num_stages)
    for stage in range(conv_stages):
        x = conv_fn(x, filters, stride_w=2, name=f"enc_s{stage + 1}", block_norm=block_norm)
        for blk in range(1, blocks_per_stage):
            x = conv_fn(x, filters, stride_w=1, name=f"enc_s{stage + 1}_b{blk}", block_norm=block_norm)
        filters = make_divisible(filters * multiplier, 8)
    for stage in range(conv_stages, num_stages):
        x = dw_fn(x, filters, stride_w=2, name=f"enc_s{stage + 1}", block_norm=block_norm)
        for blk in range(1, blocks_per_stage):
            x = dw_fn(x, filters, stride_w=1, name=f"enc_s{stage + 1}_b{blk}", block_norm=block_norm)
        filters = make_divisible(filters * multiplier, 8)
    x = keras.layers.Conv2D(embedding_dim, (1, 1), padding="same", name="to_vq")(x)
    x = _apply_norm_2d(x, head_norm, name="enc_head_norm")
    return keras.Model(inp, x, name=f"Encoder2D_ds{downsample_factor}")

compressionkit.models.rvq_autoencoder.build_decoder_2d(output_len=2048, out_ch=1, base=32, embedding_dim=16, multiplier=1.25, num_stages=4, decoder_block_norm='none', head_norm='layer', use_residual=False, activation='relu')

Build a configurable decoder that mirrors the encoder stages.

Parameters:

Name Type Description Default
output_len int

Number of output time samples.

2048
out_ch int

Number of output channels.

1
base int

Base filter count (mirroring the encoder).

32
embedding_dim int

Latent channel dimension.

16
multiplier float

Filter count multiplier per stage.

1.25
num_stages int

Number of upsample stages (must match encoder).

4
decoder_block_norm str

Normalization mode for decoder blocks.

'none'
head_norm str

Normalization mode for the output head.

'layer'
use_residual bool

If True, add shortcut connections to each stage.

False
activation str

Activation function for decoder blocks.

'relu'
Source code in compressionkit/models/decoder.py
def build_decoder_2d(
    output_len: int = 2048,
    out_ch: int = 1,
    base: int = 32,
    embedding_dim: int = 16,
    multiplier: float = 1.25,
    num_stages: int = 4,
    decoder_block_norm: str = "none",
    head_norm: str = "layer",
    use_residual: bool = False,
    activation: str = "relu",
) -> keras.Model:
    """Build a configurable decoder that mirrors the encoder stages.

    Args:
        output_len: Number of output time samples.
        out_ch: Number of output channels.
        base: Base filter count (mirroring the encoder).
        embedding_dim: Latent channel dimension.
        multiplier: Filter count multiplier per stage.
        num_stages: Number of upsample stages (must match encoder).
        decoder_block_norm: Normalization mode for decoder blocks.
        head_norm: Normalization mode for the output head.
        use_residual: If True, add shortcut connections to each stage.
        activation: Activation function for decoder blocks.
    """
    downsample_factor = 2**num_stages
    if num_stages < 1:
        raise ValueError(f"num_stages must be >= 1, got {num_stages}")
    if output_len % downsample_factor != 0:
        raise ValueError(f"output_len ({output_len}) must be divisible by 2**num_stages ({downsample_factor})")

    up_fn = res_up2d_block if use_residual else up2d_block

    inp = keras.layers.Input(
        shape=(1, output_len // downsample_factor, embedding_dim),
        name="latent_in",
    )
    x = inp
    filters = make_divisible(base * (multiplier ** max(num_stages - 1, 0)), 8)
    for stage in range(num_stages):
        x = up_fn(x, filters, name=f"dec_s{stage + 1}", block_norm=decoder_block_norm, activation=activation)
        filters = max(8, make_divisible(filters / multiplier, 8))
    x = _apply_norm_2d(x, head_norm, name="head_norm")
    out = keras.layers.Conv2D(out_ch, (1, 1), padding="same", name="out")(x)
    return keras.Model(inp, out, name=f"Decoder2D_ds{downsample_factor}")

compressionkit.models.rvq_autoencoder.compute_compression_stats(frame_size, *, bit_depth, num_channels=1, latent_width, num_levels, downsample_factor=16, bottleneck_type='rvq', fsq_levels=None, codebook_sizes=None)

Compute compression ratio and related statistics.

Supports both RVQ and FSQ bottlenecks. For FSQ, fsq_levels is required and latent_width / num_levels are ignored.

Source code in compressionkit/models/rvq_autoencoder.py
def compute_compression_stats(
    frame_size: int,
    *,
    bit_depth: int,
    num_channels: int = 1,
    latent_width: int,
    num_levels: int,
    downsample_factor: int = 16,
    bottleneck_type: str = "rvq",
    fsq_levels: list[int] | None = None,
    codebook_sizes: list[int] | None = None,
) -> dict[str, float]:
    """Compute compression ratio and related statistics.

    Supports both RVQ and FSQ bottlenecks. For FSQ, ``fsq_levels`` is
    required and ``latent_width`` / ``num_levels`` are ignored.
    """
    latent_positions = frame_size // downsample_factor
    bottleneck_type = bottleneck_type.lower()
    if bottleneck_type == "fsq":
        if not fsq_levels:
            raise ValueError("bottleneck_type='fsq' requires fsq_levels")
        codebook_size = 1
        for L in fsq_levels:
            codebook_size *= int(L)
        bits_per_index = math.log2(codebook_size)
        compressed_bits = latent_positions * bits_per_index
    elif codebook_sizes:
        bits_per_level = [math.log2(k) for k in codebook_sizes]
        bits_per_index = sum(bits_per_level) / len(bits_per_level)
        compressed_bits = latent_positions * sum(bits_per_level)
    else:
        bits_per_index = math.log2(latent_width)
        compressed_bits = latent_positions * num_levels * bits_per_index
    raw_bits = frame_size * int(num_channels) * bit_depth
    ratio = raw_bits / compressed_bits if compressed_bits else float("inf")
    result = {
        "frame_size": frame_size,
        "num_channels": int(num_channels),
        "latent_positions": latent_positions,
        "bits_per_index": bits_per_index,
        "compressed_bits_per_window": compressed_bits,
        "raw_bits_per_window": raw_bits,
        "compression_ratio": ratio,
        "bottleneck_type": bottleneck_type,
    }
    if codebook_sizes:
        result["codebook_sizes"] = codebook_sizes
        result["bits_per_level"] = [math.log2(k) for k in codebook_sizes]
    return result