Skip to content

kandinsky6_sr

Config of the Kandinsky6 SR latent-upscaler bank (latent_upscaler/config.json).

The component config is {"models": [{"target_scale": "4x" | "2x", "model": {...}}, ...], "scaling_factor": f}. Each model mapping describes one cascaded 2x+2x upsampler. Only the architecture of the released checkpoints is implemented, so every key that would select a different architecture must carry its released value; a value that is not supported raises a ValueError naming the key instead of silently building something else.

Classes

fastvideo.configs.models.upsamplers.kandinsky6_sr.Kandinsky6SRLatentUpscalerConfig dataclass

Kandinsky6SRLatentUpscalerConfig(arch_config: ArchConfig = ArchConfig(), models: list[Kandinsky6SRLatentUpscalerEntryConfig] = list(), scaling_factor: float = 1.0, scales: tuple[int, ...] = (2, 4), *, _resolved_attention_backend: AttentionBackendEnum | None = None)

Bases: UpsamplerConfig

latent_upscaler/config.json: one entry per served scale, plus the VAE latent scaling_factor.

models accepts the raw config.json entries and is normalised to :class:Kandinsky6SRLatentUpscalerEntryConfig in __post_init__ (also re-run by update_model_config).

fastvideo.configs.models.upsamplers.kandinsky6_sr.Kandinsky6SRLatentUpscalerEntryConfig dataclass

Kandinsky6SRLatentUpscalerEntryConfig(target_scale: int, in_channels: int, hidden_channels: int, stage_channels: tuple[int, int, int], num_pre_blocks: int, num_mid_blocks: int, num_post_blocks: int, expand_ratio: int, enable_x2_entry: bool = False, x2_adapter_blocks: int = 0)

One bank entry: a cascaded 2x+2x upsampler serving target_scale.

stage_channels are the widths at the 1x, 2x and 4x latent grids. With enable_x2_entry the entry also has a private x2 path (its own input stem, x2_adapter_blocks residual blocks at 1x, and private copies of the mid stage and second stage) that upsamples by 2 instead of 4.

Methods:

fastvideo.configs.models.upsamplers.kandinsky6_sr.Kandinsky6SRLatentUpscalerEntryConfig.from_dict classmethod
from_dict(spec: Any, index: int = 0) -> Kandinsky6SRLatentUpscalerEntryConfig

Validate one models[index] entry of latent_upscaler/config.json.

Source code in fastvideo/configs/models/upsamplers/kandinsky6_sr.py
@classmethod
def from_dict(cls, spec: Any, index: int = 0) -> Kandinsky6SRLatentUpscalerEntryConfig:
    """Validate one ``models[index]`` entry of ``latent_upscaler/config.json``."""
    where = f"latent upscaler models[{index}]"
    if not isinstance(spec, Mapping) or set(spec) != {"target_scale", "model"
                                                      } or not isinstance(spec["model"], Mapping):
        raise ValueError(f"{where} must be a mapping with exactly `target_scale` and a `model` mapping, "
                         f"got {spec!r}")
    target_scale = _parse_target_scale(where, spec["target_scale"])
    model = dict(spec["model"])

    if model.get("architecture") != "multi_scale":
        raise ValueError(f"{where}: `architecture` must be 'multi_scale', got {model.get('architecture')!r}")
    enable_x2_entry = model.get("enable_x2_entry", False)
    if not isinstance(enable_x2_entry, bool):
        raise ValueError(f"{where}: `enable_x2_entry` must be a boolean, got {enable_x2_entry!r}")

    fixed = dict(_FIXED_KEYS)
    if enable_x2_entry:
        fixed.update(_X2_FIXED_KEYS)
    for key, (required, omitted) in fixed.items():
        value = model.get(key, omitted)
        if value != required or type(value) is not type(required):
            raise ValueError(f"{where}: unsupported `{key}`={value!r}"
                             f"{' (omitted)' if key not in model else ''}; only {required!r} is implemented")

    x2_defaults = {
        "x2_adapter_blocks": 0,
        "x2_adapter_sources": None,
        "x2_tail_mode": "shared",
        "x2_finisher": "none"
    }
    known = {
        "architecture", "enable_x2_entry", "stage_channels", *_REQUIRED_INT_KEYS, *fixed, *x2_defaults,
        *_IGNORED_KEYS
    }
    unknown = sorted(set(model) - known)
    if unknown:
        raise ValueError(f"{where}: unknown keys {unknown}")
    if not enable_x2_entry:
        present = sorted(key for key, default in x2_defaults.items() if model.get(key, default) != default)
        if present:
            raise ValueError(f"{where}: {present} require `enable_x2_entry`=true")
    if target_scale == 2 and not enable_x2_entry:
        raise ValueError(f"{where}: `target_scale`='2x' requires `enable_x2_entry`=true")

    missing = [key for key in _REQUIRED_INT_KEYS if key not in model]
    if missing:
        raise ValueError(f"{where}: missing keys {missing}")
    ints = {key: _positive_int(where, key, model[key]) for key in _REQUIRED_INT_KEYS}
    hidden = ints["hidden_channels"]

    raw_stages = model.get("stage_channels")
    if raw_stages is None:
        stage_channels = (hidden, hidden, hidden)
    else:
        if not isinstance(raw_stages, list | tuple) or len(raw_stages) != 3:
            raise ValueError(f"{where}: `stage_channels` must list 3 widths, got {raw_stages!r}")
        first, second, third = (_positive_int(where, "stage_channels", width) for width in raw_stages)
        stage_channels = (first, second, third)
        if stage_channels[0] != hidden:
            raise ValueError(f"{where}: `hidden_channels` ({hidden}) must equal `stage_channels[0]` "
                             f"({stage_channels[0]})")

    x2_adapter_blocks = 0
    if enable_x2_entry:
        x2_adapter_blocks = _positive_int(where, "x2_adapter_blocks", model.get("x2_adapter_blocks", 0))
        # Training warm-starts the adapter from these pre_blocks; only their consistency matters here.
        sources = model.get("x2_adapter_sources")
        if sources is not None and (not isinstance(sources, list | tuple) or len(sources) != x2_adapter_blocks
                                    or any(not isinstance(i, int) or not 0 <= i < ints["num_pre_blocks"]
                                           for i in sources)):
            raise ValueError(f"{where}: `x2_adapter_sources` must list {x2_adapter_blocks} indices in "
                             f"[0, num_pre_blocks={ints['num_pre_blocks']}), got {sources!r}")

    return cls(target_scale=target_scale,
               stage_channels=stage_channels,
               enable_x2_entry=enable_x2_entry,
               x2_adapter_blocks=x2_adapter_blocks,
               **ints)