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

81 lines
2.8 KiB
Python

from __future__ import annotations
from typing import Any
import torch
from truecluster.protocol.tensors import deserialize_tensor
from truecluster.model.qwen import QwenLayerShard
class NodeRuntime:
def __init__(self, device: str):
self.device = device
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")
self.shard = QwenLayerShard(
config_dict=self.assignment["config"],
layer_start=int(self.assignment["layer_start"]),
layer_end_exclusive=int(self.assignment["layer_end_exclusive"]),
tensors=self.tensors,
device=self.device,
dtype=torch.float16 if self.device != "cpu" else torch.float32,
)
# Release CPU transfer tensors after materializing the shard.
self.tensors = {}
def clear_cache(self) -> None:
if self.shard is not None:
self.shard.clear_cache()
def forward(self, hidden: torch.Tensor, position_start: int) -> torch.Tensor:
if self.shard is None:
raise RuntimeError("model shard is not loaded")
return self.shard.forward(hidden, position_start=position_start).detach().cpu()