diff --git a/truecluster/cli.py b/truecluster/cli.py index 50dfd98..398644c 100644 --- a/truecluster/cli.py +++ b/truecluster/cli.py @@ -32,6 +32,7 @@ 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"), @@ -47,6 +48,7 @@ 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, ) @@ -62,6 +64,7 @@ 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: @@ -69,7 +72,12 @@ async def _run_cluster( log.info("loading model into coordinator RAM: %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) + 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) config = uvicorn.Config(api, host=api_host, port=api_port, log_level="info") diff --git a/truecluster/cluster/server.py b/truecluster/cluster/server.py index c9fb040..e1963a0 100644 --- a/truecluster/cluster/server.py +++ b/truecluster/cluster/server.py @@ -16,6 +16,7 @@ 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 @@ -56,6 +57,7 @@ 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 @@ -64,6 +66,7 @@ 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) @@ -129,7 +132,7 @@ class ClusterRuntime: }, ) for name, tensor in tensors.items(): - await node.send(M.WEIGHT_TENSOR, {"name": name, "tensor": serialize_tensor(tensor)}) + 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: @@ -138,6 +141,27 @@ 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: @@ -234,8 +258,8 @@ async def start_node_server(host: str, port: int, runtime: ClusterRuntime) -> as # Keep the connection open. Inference and assignment methods own reads/writes. while not reader.at_eof(): await asyncio.sleep(30) - except Exception: - log.exception("node connection failed from %s", peer) + except Exception as exc: + log.warning("node connection failed from %s: %s", peer, exc) finally: try: writer.close() diff --git a/truecluster/node/client.py b/truecluster/node/client.py index 6257b8c..03426bd 100644 --- a/truecluster/node/client.py +++ b/truecluster/node/client.py @@ -46,6 +46,10 @@ async def run_node( ) 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") runtime.load() diff --git a/truecluster/node/runtime.py b/truecluster/node/runtime.py index 44e8e2c..1aa135e 100644 --- a/truecluster/node/runtime.py +++ b/truecluster/node/runtime.py @@ -4,6 +4,8 @@ from typing import Any import torch +from truecluster.protocol.tensors import deserialize_tensor + from truecluster.model.qwen import QwenLayerShard @@ -12,16 +14,48 @@ class NodeRuntime: self.device = device 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") diff --git a/truecluster/protocol/framing.py b/truecluster/protocol/framing.py index 72e70a1..360a249 100644 --- a/truecluster/protocol/framing.py +++ b/truecluster/protocol/framing.py @@ -28,7 +28,11 @@ async def read_message(reader: asyncio.StreamReader) -> dict[str, Any] | None: (length,) = struct.unpack(">Q", header) if length > MAX_FRAME_BYTES: - raise ProtocolError(f"frame too large: {length} bytes") + preview = header.decode("ascii", errors="replace") + raise ProtocolError( + f"invalid TrueCluster frame header or frame too large: {length} bytes; " + f"first 8 bytes={preview!r}" + ) try: data = await reader.readexactly(length) except asyncio.IncompleteReadError as exc: diff --git a/truecluster/protocol/messages.py b/truecluster/protocol/messages.py index 3fc7327..7dcec13 100644 --- a/truecluster/protocol/messages.py +++ b/truecluster/protocol/messages.py @@ -1,7 +1,8 @@ HELLO = "HELLO" HELLO_ACK = "HELLO_ACK" ASSIGNMENT = "ASSIGNMENT" -WEIGHT_TENSOR = "WEIGHT_TENSOR" +WEIGHT_TENSOR = "WEIGHT_TENSOR" # legacy whole-tensor frame +WEIGHT_CHUNK = "WEIGHT_CHUNK" WEIGHTS_COMPLETE = "WEIGHTS_COMPLETE" LOAD_COMPLETE = "LOAD_COMPLETE" LOAD_FAILED = "LOAD_FAILED"