Implementing chunking
This commit is contained in:
+9
-1
@@ -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")
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user