Skip to content

minimax_h3_ref2va_dataset

Dataset contract for precomputed MiniMax H3 Ref2VA training samples.

Classes

fastvideo.dataset.minimax_h3_ref2va_dataset.MiniMaxH3Ref2VAParquetDataset

MiniMaxH3Ref2VAParquetDataset(path: str | Sequence[str] | dict[str, int], batch_size: int, *, drop_last: bool = True, seed: int = 42)

Bases: LatentsParquetMapStyleDataset

Map-style dataset for variable-length Ref2VA conditions.

Source code in fastvideo/dataset/minimax_h3_ref2va_dataset.py
def __init__(
    self,
    path: str | Sequence[str] | dict[str, int],
    batch_size: int,
    *,
    drop_last: bool = True,
    seed: int = 42,
) -> None:
    if batch_size != 1:
        raise ValueError(
            "MiniMax H3 Ref2VA requires batch_size=1"
        )

    super().__init__(
        path=path,
        batch_size=batch_size,
        parquet_schema=pyarrow_schema_minimax_h3_ref2va,
        cfg_rate=0.0,
        seed=seed,
        drop_last=drop_last,
        # __getitems__ below bypasses the generic text-padding
        # collator. The parent constructor still requires a value.
        text_padding_length=1,
    )

Functions:

fastvideo.dataset.minimax_h3_ref2va_dataset.build_minimax_h3_ref2va_dataloader

build_minimax_h3_ref2va_dataloader(path: str | Sequence[str] | dict[str, int], batch_size: int, num_data_workers: int, *, drop_last: bool = True, seed: int = 42) -> tuple[MiniMaxH3Ref2VAParquetDataset, StatefulDataLoader]

Build the stateful loader using FastVideo's DP/SP sampler.

Source code in fastvideo/dataset/minimax_h3_ref2va_dataset.py
def build_minimax_h3_ref2va_dataloader(
    path: str | Sequence[str] | dict[str, int],
    batch_size: int,
    num_data_workers: int,
    *,
    drop_last: bool = True,
    seed: int = 42,
) -> tuple[
    MiniMaxH3Ref2VAParquetDataset,
    StatefulDataLoader,
]:
    """Build the stateful loader using FastVideo's DP/SP sampler."""
    dataset = MiniMaxH3Ref2VAParquetDataset(
        path,
        batch_size,
        drop_last=drop_last,
        seed=seed,
    )

    loader = StatefulDataLoader(
        dataset,
        batch_sampler=dataset.sampler,
        collate_fn=passthrough,
        num_workers=num_data_workers,
        pin_memory=True,
        persistent_workers=num_data_workers > 0,
    )
    return dataset, loader

fastvideo.dataset.minimax_h3_ref2va_dataset.collate_minimax_h3_ref2va_rows

collate_minimax_h3_ref2va_rows(rows: list[dict[str, Any]]) -> dict[str, Any]

Collate one row without padding or truncating Qwen tokens.

Source code in fastvideo/dataset/minimax_h3_ref2va_dataset.py
def collate_minimax_h3_ref2va_rows(
    rows: list[dict[str, Any]],
) -> dict[str, Any]:
    """Collate one row without padding or truncating Qwen tokens."""
    if len(rows) != 1:
        raise ValueError(
            "MiniMax H3 Ref2VA requires exactly one row per batch, "
            f"got {len(rows)}"
        )

    row = rows[0]
    if (
        row.get("schema_version")
        != MINIMAX_H3_REF2VA_SCHEMA_VERSION
    ):
        raise ValueError(
            "Unsupported Ref2VA schema version: "
            f"{row.get('schema_version')!r}; expected "
            f"{MINIMAX_H3_REF2VA_SCHEMA_VERSION!r}"
        )

    tensors = {
        name: _decode_float32_tensor(row, name)
        for name in _TENSOR_FIELDS
    }

    text_embedding = tensors["text_embedding"]
    if (
        text_embedding.ndim != 2
        or text_embedding.shape[0] == 0
        or text_embedding.shape[1] != 5120
    ):
        raise ValueError(
            "text_embedding must have shape [length, 5120], "
            f"got {tuple(text_embedding.shape)}"
        )

    raw_tags = row.get("text_token_tags")
    if not isinstance(raw_tags, list):
        raise TypeError("text_token_tags must be a list")
    if any(
        isinstance(tag, bool) or not isinstance(tag, int)
        for tag in raw_tags
    ):
        raise TypeError("text_token_tags must contain integers")

    text_token_tags = torch.tensor(
        raw_tags,
        dtype=torch.long,
    )
    if text_token_tags.shape != text_embedding.shape[:1]:
        raise ValueError(
            "text_token_tags must align one-to-one with "
            "text_embedding"
        )
    if not bool(
        (
            (text_token_tags == 0)
            | (text_token_tags == 1)
        ).all()
    ):
        raise ValueError(
            "text_token_tags may contain only MiniMax H3 "
            "vision=0 and text=1 tags"
        )

    _validate_reference_contract(
        row,
        tensors["ref_visual_anchor"],
        tensors["ref_audio_anchor"],
    )

    info = {
        field: row.get(field)
        for field in _INFO_FIELDS
    }
    info["prompt"] = info.get("caption", "")

    return {
        **{
            name: tensor.unsqueeze(0)
            for name, tensor in tensors.items()
        },
        "text_attention_mask": torch.ones(
            (1, text_embedding.shape[0]),
            dtype=torch.float32,
        ),
        "text_token_tags": text_token_tags.unsqueeze(0),
        "info_list": [info],
        "caption_text": [info.get("caption", "")],
    }