247 lines
10 KiB
Python
247 lines
10 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__)
|
|
|
|
|
|
@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
|