260 lines
9.1 KiB
Python
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)
|