CoolFace
Modelpublic

Cccccz/HY

sourceHugging Faceupdated 11d agoView on Hugging Face
0likes
text_kv_schema.py57 linesDownload Raw Back to predictor_data
1"""Schema for exact all-layer BF16 Teacher Text K/V caches."""2 3from __future__ import annotations4 5from typing import Mapping6 7import torch8 9 10SCHEMA_VERSION_TEXT_KV_ALL54 = "hyworldplay_teacher_text_kv_exact_all54_bf16_v1"11TEXT_KV_LAYER_COUNT = 5412TEXT_KV_HEADS = 1613TEXT_KV_TOKENS = 90314TEXT_KV_HEAD_DIM = 12815 16 17def text_kv_key(block_id: int, kind: str) -> str:18    if not 0 <= block_id < TEXT_KV_LAYER_COUNT:19        raise ValueError(f"Invalid Text K/V block: {block_id}")20    if kind not in {"k_txt", "v_txt"}:21        raise ValueError(f"Invalid Text K/V kind: {kind}")22    return f"block_{block_id:02d}_{kind}"23 24 25def expected_text_kv_keys() -> set[str]:26    return {27        text_kv_key(block_id, kind)28        for block_id in range(TEXT_KV_LAYER_COUNT)29        for kind in ("k_txt", "v_txt")30    } | {"text_valid_mask"}31 32 33def validate_text_kv_all54(tensors: Mapping[str, torch.Tensor]) -> None:34    missing = expected_text_kv_keys().difference(tensors)35    unexpected = set(tensors).difference(expected_text_kv_keys())36    if missing or unexpected:37        raise ValueError(38            f"Text K/V key mismatch: missing={sorted(missing)}, "39            f"unexpected={sorted(unexpected)}"40        )41    expected_shape = (1, TEXT_KV_HEADS, TEXT_KV_TOKENS, TEXT_KV_HEAD_DIM)42    for block_id in range(TEXT_KV_LAYER_COUNT):43        for kind in ("k_txt", "v_txt"):44            name = text_kv_key(block_id, kind)45            value = tensors[name]46            if tuple(value.shape) != expected_shape:47                raise ValueError(f"Unexpected {name} shape: {tuple(value.shape)}")48            if value.dtype != torch.bfloat16:49                raise TypeError(f"{name} must be BF16, got {value.dtype}")50            if not torch.isfinite(value).all():51                raise ValueError(f"{name} contains non-finite values")52    mask = tensors["text_valid_mask"]53    if tuple(mask.shape) != (1, TEXT_KV_TOKENS) or mask.dtype != torch.bool:54        raise ValueError(55            f"Unexpected text_valid_mask: shape={tuple(mask.shape)}, dtype={mask.dtype}"56        )57