81 lines
2.8 KiB
Python
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()
|