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__) @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, ): 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._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: 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), estimated bytes=%.2f MB; node will download/load model locally", node.node_id, assignment.layer_start, assignment.layer_end_exclusive, 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, "total_weight_bytes": total_bytes, "weight_source": "huggingface_or_local_path", }, ) 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 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