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"),
|
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")
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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")
|
||||||
|
|||||||
@@ -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:
|
||||||
|
|||||||
@@ -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"
|
||||||
|
|||||||
Reference in New Issue
Block a user