Removing dumb idea

This commit is contained in:
2026-06-04 19:03:03 -05:00
parent 6e54d9eb5e
commit 04eafb5446
7 changed files with 184 additions and 186 deletions
+2
View File
@@ -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
```
+24 -49
View File
@@ -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.
+3 -5
View File
@@ -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,
)
)
+131 -50
View File
@@ -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]:
+6 -32
View File
@@ -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:
+7 -11
View File
@@ -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):
+11 -39
View File
@@ -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: