Files
truecluster/truecluster/cluster/server.py
T
2026-06-04 18:38:57 -05:00

273 lines
11 KiB
Python

from __future__ import annotations
import asyncio
import logging
import time
import uuid
from dataclasses import dataclass, field
from typing import Any
import torch
from truecluster.cluster.model_store import ModelStore
from truecluster.cluster.planner import LayerAssignment, plan_even_layers
from truecluster.protocol import messages as M
from truecluster.protocol.framing import make_message, read_message, write_message
from truecluster.protocol.tensors import deserialize_tensor, serialize_tensor
log = logging.getLogger(__name__)
DEFAULT_WEIGHT_CHUNK_BYTES = 4 * 1024 * 1024
@dataclass
class NodeHandle:
node_id: str
reader: asyncio.StreamReader
writer: asyncio.StreamWriter
hello: dict[str, Any]
assignment: LayerAssignment | None = None
ready: bool = False
lock: asyncio.Lock = field(default_factory=asyncio.Lock)
async def send(self, msg_type: str, payload: dict[str, Any] | None = None, request_id: str | None = None) -> None:
await write_message(self.writer, make_message(msg_type, payload, request_id=request_id))
async def recv(self) -> dict[str, Any]:
msg = await read_message(self.reader)
if msg is None:
raise ConnectionError(f"node {self.node_id} disconnected")
return msg
async def run_hidden(self, msg_type: str, hidden: torch.Tensor, position_start: int, request_id: str) -> torch.Tensor:
async with self.lock:
await self.send(
msg_type,
{"position_start": int(position_start), "hidden_state": serialize_tensor(hidden)},
request_id=request_id,
)
reply = await self.recv()
if reply.get("type") != M.HIDDEN_STATE:
raise RuntimeError(f"node {self.node_id} returned {reply.get('type')}: {reply.get('payload')}")
return deserialize_tensor(reply["payload"]["hidden_state"], device="cpu")
class ClusterRuntime:
def __init__(
self,
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
target_bytes = int(target_node_memory_gb * 1024**3) if target_node_memory_gb else None
layer_bytes = [model_store.layer_bytes(i) for i in range(model_store.num_layers)]
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)
@property
def ready_nodes(self) -> list[NodeHandle]:
return [n for n in self.nodes if n.ready]
@property
def is_ready(self) -> bool:
return len(self.ready_nodes) >= self.required_nodes
def readiness_error(self) -> str:
return f"Model is not ready. Required nodes: {self.required_nodes}, connected ready nodes: {len(self.ready_nodes)}"
async def add_node(self, node: NodeHandle) -> None:
self.nodes.append(node)
log.info("node connected: %s", node.node_id)
await self._maybe_assign_nodes()
async def _maybe_assign_nodes(self) -> None:
if self._assigning or self.is_ready:
return
unassigned = [n for n in self.nodes if n.assignment is None]
if len(unassigned) < self.required_nodes:
log.info("waiting for nodes: %s/%s connected", len(unassigned), self.required_nodes)
return
self._assigning = True
selected = unassigned[: self.required_nodes]
tasks = []
for node, assignment in zip(selected, self.assignments):
node.assignment = assignment
tasks.append(asyncio.create_task(self._assign_node(node, assignment)))
try:
await asyncio.gather(*tasks)
finally:
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())
log.info(
"assigning node %s layers [%s,%s), tensors=%s, bytes=%.2f MB",
node.node_id,
assignment.layer_start,
assignment.layer_end_exclusive,
len(tensors),
total_bytes / 1024 / 1024,
)
async with node.lock:
await node.send(
M.ASSIGNMENT,
{
"model_id": self.model_store.model_id,
"architecture": "qwen",
"quant": self.model_store.quant,
"compute_dtype": "fp16",
"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,
},
)
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
log.info("node ready: %s", node.node_id)
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:
await node.send(M.CLEAR_CACHE, {})
async def run_pipeline(self, hidden: torch.Tensor, position_start: int, prefill: bool, request_id: str) -> torch.Tensor:
msg_type = M.RUN_PREFILL if prefill else M.RUN_DECODE
for node in self.ready_nodes[: self.required_nodes]:
hidden = await node.run_hidden(msg_type, hidden, position_start, request_id)
return hidden
async def generate_completion(
self,
prompt: str,
max_tokens: int = 64,
temperature: float = 0.7,
top_p: float = 0.95,
stop: str | list[str] | None = None,
) -> dict[str, Any]:
from truecluster.cluster.sampler import sample_next_token
if not self.is_ready:
raise RuntimeError(self.readiness_error())
async with self.generation_lock:
await self.clear_caches()
request_id = f"req-{uuid.uuid4().hex}"
tokenizer = self.model_store.tokenizer
encoded = tokenizer(prompt, return_tensors="pt", add_special_tokens=True)
input_ids = encoded["input_ids"].to(torch.long)
prompt_tokens = int(input_ids.shape[1])
hidden = self.model_store.head.embed(input_ids)
hidden = await self.run_pipeline(hidden, position_start=0, prefill=True, request_id=request_id)
logits = self.model_store.head.logits(hidden[:, -1:, :])
generated: list[int] = []
eos_id = tokenizer.eos_token_id
stop_list = [stop] if isinstance(stop, str) else (stop or [])
finish_reason = "length"
text = ""
for step in range(max_tokens):
token = sample_next_token(logits, temperature=temperature, top_p=top_p)
generated.append(token)
text = tokenizer.decode(generated, skip_special_tokens=True)
if eos_id is not None and token == eos_id:
finish_reason = "stop"
break
if any(s and s in text for s in stop_list):
finish_reason = "stop"
break
if step == max_tokens - 1:
break
next_ids = torch.tensor([[token]], dtype=torch.long)
hidden = self.model_store.head.embed(next_ids)
hidden = await self.run_pipeline(
hidden,
position_start=prompt_tokens + step,
prefill=False,
request_id=request_id,
)
logits = self.model_store.head.logits(hidden)
# Trim at first stop string for OpenAI-like behavior.
for s in stop_list:
if s and s in text:
text = text.split(s, 1)[0]
break
return {
"text": text,
"prompt_tokens": prompt_tokens,
"completion_tokens": len(generated),
"total_tokens": prompt_tokens + len(generated),
"finish_reason": finish_reason,
}
async def start_node_server(host: str, port: int, runtime: ClusterRuntime) -> asyncio.AbstractServer:
async def handle_client(reader: asyncio.StreamReader, writer: asyncio.StreamWriter) -> None:
peer = writer.get_extra_info("peername")
try:
hello = await read_message(reader)
if not hello or hello.get("type") != M.HELLO:
await write_message(writer, make_message(M.ERROR, {"message": "expected HELLO"}))
writer.close()
await writer.wait_closed()
return
payload = hello.get("payload", {})
node_id = payload.get("node_id") or f"node-{uuid.uuid4().hex[:8]}"
node = NodeHandle(node_id=node_id, reader=reader, writer=writer, hello=payload)
await node.send(M.HELLO_ACK, {"required_nodes": runtime.required_nodes})
await runtime.add_node(node)
# Keep the connection open. Inference and assignment methods own reads/writes.
while not reader.at_eof():
await asyncio.sleep(30)
except Exception as exc:
log.warning("node connection failed from %s: %s", peer, exc)
finally:
try:
writer.close()
await writer.wait_closed()
except Exception:
pass
server = await asyncio.start_server(handle_client, host, port)
log.info("node socket server listening on %s:%s", host, port)
return server