Cccccz/HY
0
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 