diff --git a/README.md b/README.md index 0b6d3dd..569daa2 100644 --- a/README.md +++ b/README.md @@ -18,6 +18,8 @@ truecluster cluster --model Qwen/Qwen2.5-0.5B-Instruct --max-nodes 1 --quant fp1 ## Run a node +Nodes download/resolve the cluster model from HuggingFace themselves and load only the assigned layer range. + ```bash truecluster node --cluster-host 127.0.0.1 --cluster-port 7001 --device auto ``` diff --git a/SPEC.md b/SPEC.md index d8387a8..4840378 100644 --- a/SPEC.md +++ b/SPEC.md @@ -2,7 +2,7 @@ ## 1. Goal -TrueCluster is a Python 3.10 prototype for distributed LLM inference across a small cluster of heterogeneous machines. It allows a coordinator/cluster process to load a HuggingFace safetensors model, split transformer layers across connected worker nodes, send each node only the weights it needs over a socket connection, and expose a simple OpenAI-compatible HTTP API for generation. +TrueCluster is a Python 3.10 prototype for distributed LLM inference across a small cluster of heterogeneous machines. It allows a coordinator/cluster process to load a HuggingFace safetensors model config/tokenizer/coordinator-owned tensors, split transformer layers across connected worker nodes, tell each node which layer range to load from HuggingFace, and expose a simple OpenAI-compatible HTTP API for generation. The initial target is small Qwen2/Qwen2.5-style decoder-only models around 0.5B-1.5B parameters, with first-class support for mixed Nvidia CUDA and Apple Silicon Mac MPS nodes in the same cluster. Qwen3.5 hybrid linear-attention models are a later target and are not part of the first fp16 prototype. @@ -12,8 +12,8 @@ Primary prototype goals: - Easy CLI for running a cluster or node. - Coordinator loads one model at a time. - Nodes connect to the coordinator by host/port. -- Coordinator sends assigned model shards to nodes over sockets. -- Nodes do not need local model files. +- Coordinator sends layer assignments to nodes over sockets. +- Nodes download/resolve the full HuggingFace model themselves, then load only their assigned layers. - Support HuggingFace safetensors models. - Initial architecture target: Qwen2/Qwen2.5-style causal language models. - Support `fp16`, portable `int8`, and portable `int4` weight-only quantization modes. @@ -31,7 +31,7 @@ The first prototype will intentionally avoid several advanced features: - No expert parallelism. - No dynamic model hot-swapping. - No high-performance CUDA-only quantization kernels as a requirement. -- No dependency on every node having HuggingFace access. +- No coordinator-to-node weight transfer. Nodes are expected to have HuggingFace/model access. - No web UI. - No streaming responses in the first milestone. - No fault-tolerant recovery during an active generation. @@ -49,7 +49,7 @@ The coordinator owns: - Tokenizer. - Sampling logic. - HuggingFace model metadata/config loading. -- Model weight loading from safetensors. +- Coordinator-owned tensor loading from safetensors. - Model split planning. - Embedding layer. - Final normalization. @@ -81,7 +81,7 @@ HTTP request -> repeat until completion ``` -Weights are transferred once during assignment. During generation, only hidden states and small metadata are passed between coordinator and nodes. +Weights are not transferred over the cluster socket. During assignment, the coordinator sends model id, config, and layer range. Each node downloads/resolves the model from HuggingFace or a matching local path and loads only its assigned layer tensors. During generation, only hidden states and small metadata are passed between coordinator and nodes. ## 4. Execution Model @@ -173,12 +173,12 @@ Behavior: 1. Resolve/download model. 2. Load config and tokenizer. -3. Load safetensors into coordinator RAM or memory-mapped index. -4. Build model tensor index. -5. Start node socket server. -6. Start HTTP API. -7. Wait for enough nodes. -8. Assign layer ranges and transmit weights. +3. Build safetensors metadata index and load only coordinator-owned tensors. +4. Start node socket server. +5. Start HTTP API. +6. Wait for enough nodes. +7. Assign layer ranges. +8. Nodes download/resolve the model and load assigned layers. 9. Mark cluster ready. ### 5.2 Node Command @@ -199,7 +199,8 @@ Options: --cluster-port INT Coordinator node socket port. --device TEXT auto, cuda, cuda:0, mps, or cpu. Default: auto --node-id TEXT Optional stable node id. ---work-dir PATH Temporary local directory for received weights/cache. +--hf-cache-dir PATH Optional HuggingFace cache directory for node downloads. +--work-dir PATH Reserved for future temporary/cache files. --max-memory-gb FLOAT Optional memory capability override. --log-level TEXT Default: info ``` @@ -210,8 +211,8 @@ Behavior: 2. Connect to coordinator socket. 3. Send `HELLO` capability message. 4. Wait for assignment. -5. Receive config and layer weights. -6. Build local model shard. +5. Receive config and assigned layer range. +6. Download/resolve the model locally and build local model shard. 7. Mark itself ready. 8. Serve prefill/decode requests over the persistent connection. @@ -350,8 +351,6 @@ Coordinator/node lifecycle: HELLO HELLO_ACK ASSIGNMENT -WEIGHT_CHUNK -WEIGHTS_COMPLETE LOAD_COMPLETE LOAD_FAILED PING @@ -432,32 +431,9 @@ Sent by coordinator after planning. } ``` -### 7.5 Tensor Transfer +### 7.5 Model Loading on Nodes -Each tensor chunk message includes metadata: - -```json -{ - "type": "WEIGHT_CHUNK", - "payload": { - "tensor_name": "model.layers.0.self_attn.q_proj.weight", - "dtype": "float16", - "shape": [1024, 1024], - "chunk_index": 0, - "chunk_count": 4, - "offset": 0, - "data": "binary payload or external raw section" - } -} -``` - -For very large tensors, the preferred format is: - -```text -[frame length][msgpack metadata][raw bytes referenced by metadata] -``` - -The implementation should keep this hidden behind `protocol/tensors.py`. +The coordinator does not transfer model weights. After `ASSIGNMENT`, each node resolves/downloads `model_id` itself using HuggingFace cache semantics and loads only tensors whose names match its assigned layer range. ### 7.6 Inference Messages @@ -666,8 +642,7 @@ This is designed for correctness and portability, not peak speed. Preferred prototype behavior: - Coordinator loads original safetensors. -- Coordinator quantizes tensors before sending to nodes if `int8` or `int4` is selected. -- Nodes receive already-quantized tensors plus quantization metadata. +- Nodes quantize their locally loaded assigned tensors if `int8` or `int4` is selected. - Coordinator also quantizes/loads its own embedding/lm_head as needed. ## 11. Generation Algorithm @@ -907,7 +882,7 @@ Deliverables: Deliverables: - Socket protocol. -- Coordinator sends all transformer layers to one local node. +- Coordinator assigns all transformer layers to one local node. - Node loads layers and runs them. - Coordinator keeps embedding/final norm/lm head. - `/v1/completions` works. @@ -938,7 +913,7 @@ Deliverables: - `fp16` baseline. - Portable `int8` linear. - Portable `int4` linear. -- Quantized weight transfer. +- Node-local quantized layer loading. - CLI `--quant` option. ### Phase 6: API Polish @@ -957,8 +932,8 @@ A prototype is considered working when: 1. A cluster can be started with a HuggingFace safetensors Qwen-style model. 2. A node can connect to the cluster with no local model files. -3. The cluster sends layer weights to the node over the socket. -4. The node loads assigned layers on CUDA, MPS, or CPU. +3. The cluster sends a layer assignment to the node over the socket. +4. The node downloads/resolves the model and loads assigned layers on CUDA, MPS, or CPU. 5. `/v1/models` returns the loaded model. 6. `/v1/completions` generates text through the distributed pipeline. 7. `/v1/chat/completions` works for simple chat prompts. @@ -971,7 +946,7 @@ A prototype is considered working when: - Use pipeline parallelism, not tensor parallelism. - Use contiguous layer ranges only. - Keep tokenizer, embeddings, final norm, lm head, and sampler on the coordinator. -- Send weights once at node assignment time. +- Do not send weights over the node socket; nodes load weights locally from HuggingFace/cache. - Send hidden states during generation. - Store KV cache on worker nodes. - Implement portable quantization instead of relying on CUDA-only libraries. diff --git a/truecluster/cli.py b/truecluster/cli.py index 398644c..aa27eca 100644 --- a/truecluster/cli.py +++ b/truecluster/cli.py @@ -32,7 +32,6 @@ def cluster_cmd( max_nodes: int = typer.Option(1, "--max-nodes", min=1, help="Maximum/requested worker shard count"), quant: str = typer.Option("fp16", "--quant", help="Weight mode; currently only fp16 is implemented"), target_node_memory_gb: Optional[float] = typer.Option(None, "--target-node-memory-gb", help="Optional planner memory hint"), - weight_chunk_mb: float = typer.Option(4.0, "--weight-chunk-mb", min=0.0625, help="Model weight transfer chunk size in MiB"), trust_remote_code: bool = typer.Option(False, "--trust-remote-code", help="Allow HF remote code for config/tokenizer"), hf_cache_dir: Optional[str] = typer.Option(None, "--hf-cache-dir", help="Optional HuggingFace cache directory"), log_level: str = typer.Option("info", "--log-level"), @@ -48,7 +47,6 @@ def cluster_cmd( max_nodes=max_nodes, quant=quant, target_node_memory_gb=target_node_memory_gb, - weight_chunk_mb=weight_chunk_mb, trust_remote_code=trust_remote_code, hf_cache_dir=hf_cache_dir, ) @@ -64,19 +62,17 @@ async def _run_cluster( max_nodes: int, quant: str, target_node_memory_gb: float | None, - weight_chunk_mb: float, trust_remote_code: bool, hf_cache_dir: str | None, ) -> None: log = logging.getLogger("truecluster.cluster") - log.info("loading model into coordinator RAM: %s", model) + log.info("loading coordinator model metadata/tokenizer/head tensors: %s", model) store = ModelStore.load(model, quant=quant, trust_remote_code=trust_remote_code, hf_cache_dir=hf_cache_dir) log.info("model loaded: %s layers, %s tensors", store.num_layers, len(store.tensors)) runtime = ClusterRuntime( store, max_nodes=max_nodes, target_node_memory_gb=target_node_memory_gb, - weight_chunk_mb=weight_chunk_mb, ) node_server = await start_node_server(node_host, node_port, runtime) api = create_app(runtime) @@ -95,6 +91,7 @@ def node_cmd( node_id: Optional[str] = typer.Option(None, "--node-id", help="Optional stable node id"), work_dir: Optional[str] = typer.Option(None, "--work-dir", help="Reserved for future temporary weight storage"), max_memory_gb: Optional[float] = typer.Option(None, "--max-memory-gb", help="Optional capability override"), + hf_cache_dir: Optional[str] = typer.Option(None, "--hf-cache-dir", help="Optional HuggingFace cache directory for node downloads"), log_level: str = typer.Option("info", "--log-level"), ) -> None: _setup_logging(log_level) @@ -107,6 +104,7 @@ def node_cmd( device=device, node_id=node_id, max_memory_gb=max_memory_gb, + hf_cache_dir=hf_cache_dir, ) ) diff --git a/truecluster/cluster/model_store.py b/truecluster/cluster/model_store.py index 73d3a10..943bb9c 100644 --- a/truecluster/cluster/model_store.py +++ b/truecluster/cluster/model_store.py @@ -1,12 +1,13 @@ 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.torch import load_file +from safetensors import safe_open from transformers import AutoConfig, AutoTokenizer from truecluster.model.qwen import CoordinatorHead, QwenConfig @@ -16,11 +17,39 @@ 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 the HF/local model, loads safetensors into RAM, and - keeps that RAM copy available for fast assignment transfer to worker nodes. + 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__( @@ -30,6 +59,7 @@ class ModelStore: 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 @@ -37,6 +67,7 @@ class ModelStore: 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) @@ -52,59 +83,60 @@ class ModelStore: if quant != "fp16": raise ModelStoreError("only fp16 is implemented currently; int8/int4 are planned") - path = Path(model).expanduser() - if path.exists(): - model_path = path.resolve() - model_id = model - else: - model_path = Path( - snapshot_download( - repo_id=model, - cache_dir=hf_cache_dir, - allow_patterns=[ - "*.json", - "*.safetensors", - "*.model", - "tokenizer*", - "vocab*", - "merges.txt", - ], - ) - ) - model_id = model - + 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) - tensors = _load_safetensors_into_ram(model_path) + 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) - _validate_qwen_tensors(config_dict, tensors) - return cls(model_id, model_path, config_dict, tokenizer, tensors, quant=quant) + 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]: - out: dict[str, torch.Tensor] = {} - prefixes = [f"model.layers.{i}." for i in range(layer_start, layer_end_exclusive)] - for name, tensor in self.tensors.items(): - if any(name.startswith(prefix) for prefix in prefixes): - out[name] = tensor - return out + 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(t.numel() * t.element_size() for name, t in self.tensors.items() if name.startswith(prefix)) + 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. - - The current runtime implements Qwen2/Qwen2.5-style full self-attention - decoder blocks. It intentionally does not implement Qwen3.5's hybrid - GatedDeltaNet/linear-attention blocks yet. - """ + """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: @@ -127,21 +159,67 @@ def _normalize_supported_config(config_dict: dict[str, Any]) -> dict[str, Any]: return config_dict -def _load_safetensors_into_ram(model_path: Path) -> dict[str, torch.Tensor]: +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}") - tensors: dict[str, torch.Tensor] = {} + index: dict[str, TensorInfo] = {} for file in files: - part = load_file(str(file), device="cpu") - overlap = set(tensors).intersection(part) - if overlap: - raise ModelStoreError(f"duplicate tensor names in safetensors: {sorted(overlap)[:5]}") - tensors.update(part) + 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: @@ -150,7 +228,7 @@ def _convert_float_tensors_to_fp16(tensors: dict[str, torch.Tensor]) -> None: tensors[name] = tensor.contiguous() -def _validate_qwen_tensors(config_dict: dict[str, Any], tensors: dict[str, torch.Tensor]) -> None: +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): @@ -168,9 +246,12 @@ def _validate_qwen_tensors(config_dict: dict[str, Any], tensors: dict[str, torch f"{p}.mlp.down_proj.weight", ] ) - missing = [name for name in required if name not in tensors] + missing = [name for name in required if name not in tensor_names] if missing: - raise ModelStoreError("model does not look like supported Qwen-style safetensors; missing: " + ", ".join(missing[:20])) + 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]: diff --git a/truecluster/cluster/server.py b/truecluster/cluster/server.py index e1963a0..61f5451 100644 --- a/truecluster/cluster/server.py +++ b/truecluster/cluster/server.py @@ -16,7 +16,6 @@ from truecluster.protocol.framing import make_message, read_message, write_messa from truecluster.protocol.tensors import deserialize_tensor, serialize_tensor log = logging.getLogger(__name__) -DEFAULT_WEIGHT_CHUNK_BYTES = 4 * 1024 * 1024 @dataclass @@ -57,7 +56,6 @@ class ClusterRuntime: model_store: ModelStore, max_nodes: int, target_node_memory_gb: float | None = None, - weight_chunk_mb: float = 4.0, ): self.model_store = model_store self.max_nodes = max_nodes @@ -66,7 +64,6 @@ class ClusterRuntime: self.assignments = plan_even_layers(model_store.num_layers, max_nodes, target_bytes, layer_bytes) self.required_nodes = len(self.assignments) self.nodes: list[NodeHandle] = [] - self.weight_chunk_bytes = max(64 * 1024, int(weight_chunk_mb * 1024 * 1024)) self._assigning = False self.generation_lock = asyncio.Lock() log.info("planned %s assignment(s): %s", self.required_nodes, self.assignments) @@ -106,14 +103,15 @@ class ClusterRuntime: self._assigning = False async def _assign_node(self, node: NodeHandle, assignment: LayerAssignment) -> None: - tensors = self.model_store.tensors_for_layers(assignment.layer_start, assignment.layer_end_exclusive) - total_bytes = sum(t.numel() * t.element_size() for t in tensors.values()) + total_bytes = sum( + self.model_store.layer_bytes(i) + for i in range(assignment.layer_start, assignment.layer_end_exclusive) + ) log.info( - "assigning node %s layers [%s,%s), tensors=%s, bytes=%.2f MB", + "assigning node %s layers [%s,%s), estimated bytes=%.2f MB; node will download/load model locally", node.node_id, assignment.layer_start, assignment.layer_end_exclusive, - len(tensors), total_bytes / 1024 / 1024, ) async with node.lock: @@ -127,13 +125,10 @@ class ClusterRuntime: "layer_start": assignment.layer_start, "layer_end_exclusive": assignment.layer_end_exclusive, "config": self.model_store.config_dict, - "tensor_count": len(tensors), "total_weight_bytes": total_bytes, + "weight_source": "huggingface_or_local_path", }, ) - for name, tensor in tensors.items(): - await self._send_weight_tensor(node, name, tensor) - await node.send(M.WEIGHTS_COMPLETE, {}) reply = await node.recv() if reply.get("type") == M.LOAD_COMPLETE: node.ready = True @@ -141,27 +136,6 @@ class ClusterRuntime: return raise RuntimeError(f"node {node.node_id} failed to load: {reply}") - async def _send_weight_tensor(self, node: NodeHandle, name: str, tensor: torch.Tensor) -> None: - serialized = serialize_tensor(tensor) - data = serialized["data"] - total_bytes = len(data) - chunk_count = max(1, (total_bytes + self.weight_chunk_bytes - 1) // self.weight_chunk_bytes) - for chunk_index in range(chunk_count): - start = chunk_index * self.weight_chunk_bytes - end = min(start + self.weight_chunk_bytes, total_bytes) - await node.send( - M.WEIGHT_CHUNK, - { - "name": name, - "dtype": serialized["dtype"], - "shape": serialized["shape"], - "total_bytes": total_bytes, - "chunk_index": chunk_index, - "chunk_count": chunk_count, - "data": data[start:end], - }, - ) - async def clear_caches(self) -> None: for node in self.ready_nodes[: self.required_nodes]: async with node.lock: diff --git a/truecluster/node/client.py b/truecluster/node/client.py index 03426bd..1ffbe56 100644 --- a/truecluster/node/client.py +++ b/truecluster/node/client.py @@ -18,9 +18,10 @@ async def run_node( device: str = "auto", node_id: str | None = None, max_memory_gb: float | None = None, + hf_cache_dir: str | None = None, ) -> None: selected = select_device(device) - runtime = NodeRuntime(selected) + runtime = NodeRuntime(selected, hf_cache_dir=hf_cache_dir) log.info("connecting to cluster %s:%s using device %s", cluster_host, cluster_port, selected) reader, writer = await asyncio.open_connection(cluster_host, cluster_port) await write_message(writer, make_message(M.HELLO, capabilities(selected, node_id=node_id, max_memory_gb=max_memory_gb))) @@ -38,23 +39,18 @@ async def run_node( elif msg_type == M.ASSIGNMENT: runtime.set_assignment(payload) log.info( - "received assignment: layers [%s,%s), tensors=%s, bytes=%.2f MB", + "received assignment: model=%s layers [%s,%s), estimated bytes=%.2f MB", + payload.get("model_id"), payload.get("layer_start"), payload.get("layer_end_exclusive"), - payload.get("tensor_count"), int(payload.get("total_weight_bytes", 0)) / 1024 / 1024, ) - elif msg_type == M.WEIGHT_TENSOR: - runtime.add_tensor(payload["name"], deserialize_tensor(payload["tensor"], device="cpu")) - elif msg_type == M.WEIGHT_CHUNK: - completed = runtime.add_tensor_chunk(payload) - if completed and len(runtime.tensors) % 25 == 0: - log.info("received %s tensors", len(runtime.tensors)) - elif msg_type == M.WEIGHTS_COMPLETE: - log.info("all weights received; loading shard") + log.info("downloading/resolving model locally and loading assigned layers") runtime.load() await write_message(writer, make_message(M.LOAD_COMPLETE, {"device": selected})) log.info("shard loaded and ready") + elif msg_type in (M.WEIGHT_TENSOR, M.WEIGHT_CHUNK, M.WEIGHTS_COMPLETE): + log.warning("cluster tried to send weights, but this protocol now expects local node downloads") elif msg_type == M.CLEAR_CACHE: runtime.clear_cache() elif msg_type in (M.RUN_PREFILL, M.RUN_DECODE): diff --git a/truecluster/node/runtime.py b/truecluster/node/runtime.py index 1aa135e..e1da4a9 100644 --- a/truecluster/node/runtime.py +++ b/truecluster/node/runtime.py @@ -4,61 +4,33 @@ from typing import Any import torch -from truecluster.protocol.tensors import deserialize_tensor - +from truecluster.cluster.model_store import load_layer_tensors_for_node from truecluster.model.qwen import QwenLayerShard class NodeRuntime: - def __init__(self, device: str): + def __init__(self, device: str, hf_cache_dir: str | None = None): self.device = device + self.hf_cache_dir = hf_cache_dir self.assignment: dict[str, Any] | None = None self.tensors: dict[str, torch.Tensor] = {} - self._tensor_chunks: dict[str, dict[str, Any]] = {} self.shard: QwenLayerShard | None = None def set_assignment(self, payload: dict[str, Any]) -> None: self.assignment = payload self.tensors = {} - self._tensor_chunks = {} self.shard = None - def add_tensor(self, name: str, tensor: torch.Tensor) -> None: - self.tensors[name] = tensor - - def add_tensor_chunk(self, payload: dict[str, Any]) -> bool: - """Add one incoming weight chunk. - - Returns True when the full tensor has been reassembled. - """ - - name = payload["name"] - entry = self._tensor_chunks.setdefault( - name, - { - "dtype": payload["dtype"], - "shape": payload["shape"], - "total_bytes": int(payload["total_bytes"]), - "chunk_count": int(payload["chunk_count"]), - "chunks": {}, - }, - ) - entry["chunks"][int(payload["chunk_index"])] = payload["data"] - if len(entry["chunks"]) != entry["chunk_count"]: - return False - data = b"".join(entry["chunks"][i] for i in range(entry["chunk_count"])) - if len(data) != entry["total_bytes"]: - raise RuntimeError(f"reassembled tensor {name} has {len(data)} bytes, expected {entry['total_bytes']}") - self.tensors[name] = deserialize_tensor( - {"dtype": entry["dtype"], "shape": entry["shape"], "data": data}, - device="cpu", - ) - del self._tensor_chunks[name] - return True - def load(self) -> None: if self.assignment is None: raise RuntimeError("no assignment received") + if not self.tensors: + self.tensors = load_layer_tensors_for_node( + self.assignment["model_id"], + int(self.assignment["layer_start"]), + int(self.assignment["layer_end_exclusive"]), + hf_cache_dir=self.hf_cache_dir, + ) self.shard = QwenLayerShard( config_dict=self.assignment["config"], layer_start=int(self.assignment["layer_start"]), @@ -67,7 +39,7 @@ class NodeRuntime: device=self.device, dtype=torch.float16 if self.device != "cpu" else torch.float32, ) - # Release CPU transfer tensors after materializing the shard. + # Release CPU tensors after materializing the shard on the selected device. self.tensors = {} def clear_cache(self) -> None: