Files
truecluster/truecluster/cluster/model_store.py
T
2026-06-04 19:03:03 -05:00

260 lines
9.1 KiB
Python

from __future__ import annotations
import json
from dataclasses import dataclass
from pathlib import Path
from typing import Any
import torch
from huggingface_hub import snapshot_download
from safetensors import safe_open
from transformers import AutoConfig, AutoTokenizer
from truecluster.model.qwen import CoordinatorHead, QwenConfig
class ModelStoreError(RuntimeError):
pass
_DTYPE_BYTES = {
"BOOL": 1,
"U8": 1,
"I8": 1,
"I16": 2,
"U16": 2,
"F16": 2,
"BF16": 2,
"I32": 4,
"U32": 4,
"F32": 4,
"F64": 8,
"I64": 8,
"U64": 8,
}
@dataclass(frozen=True)
class TensorInfo:
name: str
file: Path
shape: list[int]
dtype: str
nbytes: int
class ModelStore:
"""Coordinator-side model owner.
The coordinator resolves/downloads the model, loads tokenizer/config, and
only loads the small coordinator-owned tensors into RAM: embeddings, final
norm, and lm_head if present. Worker nodes download the same HuggingFace
model themselves and load only their assigned layer range.
"""
def __init__(
self,
model_id: str,
model_path: Path,
config_dict: dict[str, Any],
tokenizer: Any,
tensors: dict[str, torch.Tensor],
tensor_index: dict[str, TensorInfo],
quant: str = "fp16",
):
self.model_id = model_id
self.model_path = model_path
self.config_dict = config_dict
self.tokenizer = tokenizer
self.tensors = tensors
self.tensor_index = tensor_index
self.quant = quant
self.qwen_config = QwenConfig.from_dict(config_dict)
self.head = CoordinatorHead(config_dict, tensors, dtype=torch.float32)
@classmethod
def load(
cls,
model: str,
quant: str = "fp16",
trust_remote_code: bool = False,
hf_cache_dir: str | None = None,
) -> "ModelStore":
if quant != "fp16":
raise ModelStoreError("only fp16 is implemented currently; int8/int4 are planned")
model_path, model_id = resolve_model_path(model, hf_cache_dir=hf_cache_dir)
config = AutoConfig.from_pretrained(str(model_path), trust_remote_code=trust_remote_code)
config_dict = _normalize_supported_config(config.to_dict())
tokenizer = AutoTokenizer.from_pretrained(str(model_path), trust_remote_code=trust_remote_code)
tensor_index = build_tensor_index(model_path)
_validate_qwen_tensor_names(config_dict, set(tensor_index))
head_names = ["model.embed_tokens.weight", "model.norm.weight"]
if "lm_head.weight" in tensor_index:
head_names.append("lm_head.weight")
tensors = load_selected_tensors(model_path, head_names)
_convert_float_tensors_to_fp16(tensors)
return cls(model_id, model_path, config_dict, tokenizer, tensors, tensor_index, quant=quant)
@property
def num_layers(self) -> int:
return self.qwen_config.num_hidden_layers
def tensors_for_layers(self, layer_start: int, layer_end_exclusive: int) -> dict[str, torch.Tensor]:
names = layer_tensor_names_from_index(self.tensor_index, layer_start, layer_end_exclusive)
tensors = load_selected_tensors(self.model_path, names)
_convert_float_tensors_to_fp16(tensors)
return tensors
def layer_bytes(self, layer_idx: int) -> int:
prefix = f"model.layers.{layer_idx}."
return sum(info.nbytes for name, info in self.tensor_index.items() if name.startswith(prefix))
def resolve_model_path(model: str, hf_cache_dir: str | None = None) -> tuple[Path, str]:
path = Path(model).expanduser()
if path.exists():
return path.resolve(), model
return (
Path(
snapshot_download(
repo_id=model,
cache_dir=hf_cache_dir,
allow_patterns=[
"*.json",
"*.safetensors",
"*.model",
"tokenizer*",
"vocab*",
"merges.txt",
],
)
),
model,
)
def _normalize_supported_config(config_dict: dict[str, Any]) -> dict[str, Any]:
"""Return the text decoder config for architectures this prototype supports."""
model_type = config_dict.get("model_type")
if model_type == "qwen3_5" or "text_config" in config_dict:
text_cfg = config_dict.get("text_config") or {}
layer_types = text_cfg.get("layer_types") or []
if any(layer_type != "full_attention" for layer_type in layer_types):
raise ModelStoreError(
"Qwen3.5 hybrid/linear-attention models are not supported by the fp16 prototype yet. "
"Use a Qwen2/Qwen2.5 causal LM for now, for example "
"Qwen/Qwen2.5-0.5B-Instruct or Qwen/Qwen2.5-1.5B-Instruct."
)
if text_cfg:
return text_cfg
supported = {"qwen2", "qwen2_moe", "qwen3"}
if model_type not in supported:
raise ModelStoreError(
f"unsupported model_type {model_type!r}; current prototype supports Qwen2/Qwen2.5-style causal LMs"
)
return config_dict
def build_tensor_index(model_path: Path) -> dict[str, TensorInfo]:
files = sorted(model_path.glob("*.safetensors"))
if not files:
raise ModelStoreError(f"no safetensors files found in {model_path}")
index: dict[str, TensorInfo] = {}
for file in files:
with safe_open(str(file), framework="pt", device="cpu") as f:
for name in f.keys():
if name in index:
raise ModelStoreError(f"duplicate tensor name in safetensors: {name}")
view = f.get_slice(name)
shape = list(view.get_shape())
dtype = str(view.get_dtype())
elem_size = _DTYPE_BYTES.get(dtype)
if elem_size is None:
raise ModelStoreError(f"unsupported safetensors dtype {dtype} for {name}")
nbytes = elem_size
for dim in shape:
nbytes *= int(dim)
index[name] = TensorInfo(name=name, file=file, shape=shape, dtype=dtype, nbytes=nbytes)
return index
def load_selected_tensors(model_path: Path, names: list[str]) -> dict[str, torch.Tensor]:
wanted = set(names)
tensors: dict[str, torch.Tensor] = {}
for file in sorted(model_path.glob("*.safetensors")):
with safe_open(str(file), framework="pt", device="cpu") as f:
available = wanted.intersection(f.keys())
for name in available:
tensors[name] = f.get_tensor(name).contiguous()
missing = wanted.difference(tensors)
if missing:
raise ModelStoreError(f"missing tensors: {sorted(missing)[:20]}")
return tensors
def load_layer_tensors_for_node(
model: str,
layer_start: int,
layer_end_exclusive: int,
hf_cache_dir: str | None = None,
) -> dict[str, torch.Tensor]:
"""Download/resolve the whole model and load only the assigned layer tensors."""
model_path, _ = resolve_model_path(model, hf_cache_dir=hf_cache_dir)
tensor_index = build_tensor_index(model_path)
names = layer_tensor_names_from_index(tensor_index, layer_start, layer_end_exclusive)
tensors = load_selected_tensors(model_path, names)
_convert_float_tensors_to_fp16(tensors)
return tensors
def layer_tensor_names_from_index(
tensor_index: dict[str, TensorInfo], layer_start: int, layer_end_exclusive: int
) -> list[str]:
prefixes = [f"model.layers.{i}." for i in range(layer_start, layer_end_exclusive)]
return [name for name in tensor_index if any(name.startswith(prefix) for prefix in prefixes)]
def _convert_float_tensors_to_fp16(tensors: dict[str, torch.Tensor]) -> None:
for name, tensor in list(tensors.items()):
if tensor.is_floating_point() and tensor.dtype != torch.float16:
tensors[name] = tensor.to(torch.float16).contiguous()
else:
tensors[name] = tensor.contiguous()
def _validate_qwen_tensor_names(config_dict: dict[str, Any], tensor_names: set[str]) -> None:
cfg = QwenConfig.from_dict(config_dict)
required = ["model.embed_tokens.weight", "model.norm.weight"]
for i in range(cfg.num_hidden_layers):
p = f"model.layers.{i}"
required.extend(
[
f"{p}.input_layernorm.weight",
f"{p}.self_attn.q_proj.weight",
f"{p}.self_attn.k_proj.weight",
f"{p}.self_attn.v_proj.weight",
f"{p}.self_attn.o_proj.weight",
f"{p}.post_attention_layernorm.weight",
f"{p}.mlp.gate_proj.weight",
f"{p}.mlp.up_proj.weight",
f"{p}.mlp.down_proj.weight",
]
)
missing = [name for name in required if name not in tensor_names]
if missing:
raise ModelStoreError(
"model does not look like supported Qwen2/Qwen2.5-style safetensors; missing: "
+ ", ".join(missing[:20])
)
def read_config_json(model_path: Path) -> dict[str, Any]:
with (model_path / "config.json").open("r", encoding="utf-8") as f:
return json.load(f)