Implementing chunking

This commit is contained in:
2026-06-04 18:38:57 -05:00
parent 9bd7f5838e
commit 6e54d9eb5e
6 changed files with 81 additions and 6 deletions
+9 -1
View File
@@ -32,6 +32,7 @@ def cluster_cmd(
max_nodes: int = typer.Option(1, "--max-nodes", min=1, help="Maximum/requested worker shard count"), 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"), 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"), 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"), 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"), hf_cache_dir: Optional[str] = typer.Option(None, "--hf-cache-dir", help="Optional HuggingFace cache directory"),
log_level: str = typer.Option("info", "--log-level"), log_level: str = typer.Option("info", "--log-level"),
@@ -47,6 +48,7 @@ def cluster_cmd(
max_nodes=max_nodes, max_nodes=max_nodes,
quant=quant, quant=quant,
target_node_memory_gb=target_node_memory_gb, target_node_memory_gb=target_node_memory_gb,
weight_chunk_mb=weight_chunk_mb,
trust_remote_code=trust_remote_code, trust_remote_code=trust_remote_code,
hf_cache_dir=hf_cache_dir, hf_cache_dir=hf_cache_dir,
) )
@@ -62,6 +64,7 @@ async def _run_cluster(
max_nodes: int, max_nodes: int,
quant: str, quant: str,
target_node_memory_gb: float | None, target_node_memory_gb: float | None,
weight_chunk_mb: float,
trust_remote_code: bool, trust_remote_code: bool,
hf_cache_dir: str | None, hf_cache_dir: str | None,
) -> None: ) -> None:
@@ -69,7 +72,12 @@ async def _run_cluster(
log.info("loading model into coordinator RAM: %s", model) 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) 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)) 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) node_server = await start_node_server(node_host, node_port, runtime)
api = create_app(runtime) api = create_app(runtime)
config = uvicorn.Config(api, host=api_host, port=api_port, log_level="info") config = uvicorn.Config(api, host=api_host, port=api_port, log_level="info")
+27 -3
View File
@@ -16,6 +16,7 @@ from truecluster.protocol.framing import make_message, read_message, write_messa
from truecluster.protocol.tensors import deserialize_tensor, serialize_tensor from truecluster.protocol.tensors import deserialize_tensor, serialize_tensor
log = logging.getLogger(__name__) log = logging.getLogger(__name__)
DEFAULT_WEIGHT_CHUNK_BYTES = 4 * 1024 * 1024
@dataclass @dataclass
@@ -56,6 +57,7 @@ class ClusterRuntime:
model_store: ModelStore, model_store: ModelStore,
max_nodes: int, max_nodes: int,
target_node_memory_gb: float | None = None, target_node_memory_gb: float | None = None,
weight_chunk_mb: float = 4.0,
): ):
self.model_store = model_store self.model_store = model_store
self.max_nodes = max_nodes 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.assignments = plan_even_layers(model_store.num_layers, max_nodes, target_bytes, layer_bytes)
self.required_nodes = len(self.assignments) self.required_nodes = len(self.assignments)
self.nodes: list[NodeHandle] = [] self.nodes: list[NodeHandle] = []
self.weight_chunk_bytes = max(64 * 1024, int(weight_chunk_mb * 1024 * 1024))
self._assigning = False self._assigning = False
self.generation_lock = asyncio.Lock() self.generation_lock = asyncio.Lock()
log.info("planned %s assignment(s): %s", self.required_nodes, self.assignments) log.info("planned %s assignment(s): %s", self.required_nodes, self.assignments)
@@ -129,7 +132,7 @@ class ClusterRuntime:
}, },
) )
for name, tensor in tensors.items(): 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, {}) await node.send(M.WEIGHTS_COMPLETE, {})
reply = await node.recv() reply = await node.recv()
if reply.get("type") == M.LOAD_COMPLETE: if reply.get("type") == M.LOAD_COMPLETE:
@@ -138,6 +141,27 @@ class ClusterRuntime:
return return
raise RuntimeError(f"node {node.node_id} failed to load: {reply}") 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: async def clear_caches(self) -> None:
for node in self.ready_nodes[: self.required_nodes]: for node in self.ready_nodes[: self.required_nodes]:
async with node.lock: 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. # Keep the connection open. Inference and assignment methods own reads/writes.
while not reader.at_eof(): while not reader.at_eof():
await asyncio.sleep(30) await asyncio.sleep(30)
except Exception: except Exception as exc:
log.exception("node connection failed from %s", peer) log.warning("node connection failed from %s: %s", peer, exc)
finally: finally:
try: try:
writer.close() writer.close()
+4
View File
@@ -46,6 +46,10 @@ async def run_node(
) )
elif msg_type == M.WEIGHT_TENSOR: elif msg_type == M.WEIGHT_TENSOR:
runtime.add_tensor(payload["name"], deserialize_tensor(payload["tensor"], device="cpu")) 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: elif msg_type == M.WEIGHTS_COMPLETE:
log.info("all weights received; loading shard") log.info("all weights received; loading shard")
runtime.load() runtime.load()
+34
View File
@@ -4,6 +4,8 @@ from typing import Any
import torch import torch
from truecluster.protocol.tensors import deserialize_tensor
from truecluster.model.qwen import QwenLayerShard from truecluster.model.qwen import QwenLayerShard
@@ -12,16 +14,48 @@ class NodeRuntime:
self.device = device self.device = device
self.assignment: dict[str, Any] | None = None self.assignment: dict[str, Any] | None = None
self.tensors: dict[str, torch.Tensor] = {} self.tensors: dict[str, torch.Tensor] = {}
self._tensor_chunks: dict[str, dict[str, Any]] = {}
self.shard: QwenLayerShard | None = None self.shard: QwenLayerShard | None = None
def set_assignment(self, payload: dict[str, Any]) -> None: def set_assignment(self, payload: dict[str, Any]) -> None:
self.assignment = payload self.assignment = payload
self.tensors = {} self.tensors = {}
self._tensor_chunks = {}
self.shard = None self.shard = None
def add_tensor(self, name: str, tensor: torch.Tensor) -> None: def add_tensor(self, name: str, tensor: torch.Tensor) -> None:
self.tensors[name] = tensor 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: def load(self) -> None:
if self.assignment is None: if self.assignment is None:
raise RuntimeError("no assignment received") raise RuntimeError("no assignment received")
+5 -1
View File
@@ -28,7 +28,11 @@ async def read_message(reader: asyncio.StreamReader) -> dict[str, Any] | None:
(length,) = struct.unpack(">Q", header) (length,) = struct.unpack(">Q", header)
if length > MAX_FRAME_BYTES: 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: try:
data = await reader.readexactly(length) data = await reader.readexactly(length)
except asyncio.IncompleteReadError as exc: except asyncio.IncompleteReadError as exc:
+2 -1
View File
@@ -1,7 +1,8 @@
HELLO = "HELLO" HELLO = "HELLO"
HELLO_ACK = "HELLO_ACK" HELLO_ACK = "HELLO_ACK"
ASSIGNMENT = "ASSIGNMENT" ASSIGNMENT = "ASSIGNMENT"
WEIGHT_TENSOR = "WEIGHT_TENSOR" WEIGHT_TENSOR = "WEIGHT_TENSOR" # legacy whole-tensor frame
WEIGHT_CHUNK = "WEIGHT_CHUNK"
WEIGHTS_COMPLETE = "WEIGHTS_COMPLETE" WEIGHTS_COMPLETE = "WEIGHTS_COMPLETE"
LOAD_COMPLETE = "LOAD_COMPLETE" LOAD_COMPLETE = "LOAD_COMPLETE"
LOAD_FAILED = "LOAD_FAILED" LOAD_FAILED = "LOAD_FAILED"