commit 9bd7f5838e4530d8c51ddd9de7700a5d9e3533f0 Author: Owen Qwen Date: Thu Jun 4 18:11:26 2026 -0500 Inital commit diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..f904eca --- /dev/null +++ b/.gitignore @@ -0,0 +1,47 @@ +# Python +__pycache__/ +*.py[cod] +*$py.class +*.so +.Python +.venv/ +venv/ +env/ +ENV/ +build/ +dist/ +*.egg-info/ +.eggs/ +.pytest_cache/ +.ruff_cache/ +.mypy_cache/ +.coverage +htmlcov/ + +# IDE/editor +.vscode/ +.idea/ +*.swp +*.swo +.DS_Store + +# Runtime/logs +*.log +logs/ +run/ +*.pid + +# Model/cache artifacts +models/ +checkpoints/ +*.safetensors +*.bin +*.pt +*.pth +*.gguf +huggingface/ +.cache/ + +# TrueCluster temp files +truecluster-work/ +node-work/ diff --git a/README.md b/README.md new file mode 100644 index 0000000..0b6d3dd --- /dev/null +++ b/README.md @@ -0,0 +1,31 @@ +# TrueCluster + +Prototype heterogeneous pipeline-parallel LLM inference cluster. + +See [`SPEC.md`](SPEC.md). + +## Install + +```bash +pip install -e . +``` + +## Run a cluster + +```bash +truecluster cluster --model Qwen/Qwen2.5-0.5B-Instruct --max-nodes 1 --quant fp16 +``` + +## Run a node + +```bash +truecluster node --cluster-host 127.0.0.1 --cluster-port 7001 --device auto +``` + +## Generate + +```bash +curl http://127.0.0.1:8000/v1/completions \ + -H 'content-type: application/json' \ + -d '{"model":"Qwen/Qwen2.5-0.5B-Instruct","prompt":"Hello","max_tokens":32}' +``` diff --git a/SPEC.md b/SPEC.md new file mode 100644 index 0000000..d8387a8 --- /dev/null +++ b/SPEC.md @@ -0,0 +1,979 @@ +# TrueCluster Prototype Specification + +## 1. Goal + +TrueCluster is a Python 3.10 prototype for distributed LLM inference across a small cluster of heterogeneous machines. It allows a coordinator/cluster process to load a HuggingFace safetensors model, split transformer layers across connected worker nodes, send each node only the weights it needs over a socket connection, and expose a simple OpenAI-compatible HTTP API for generation. + +The initial target is small Qwen2/Qwen2.5-style decoder-only models around 0.5B-1.5B parameters, with first-class support for mixed Nvidia CUDA and Apple Silicon Mac MPS nodes in the same cluster. Qwen3.5 hybrid linear-attention models are a later target and are not part of the first fp16 prototype. + +Primary prototype goals: + +- Python 3.10. +- Easy CLI for running a cluster or node. +- Coordinator loads one model at a time. +- Nodes connect to the coordinator by host/port. +- Coordinator sends assigned model shards to nodes over sockets. +- Nodes do not need local model files. +- Support HuggingFace safetensors models. +- Initial architecture target: Qwen2/Qwen2.5-style causal language models. +- Support `fp16`, portable `int8`, and portable `int4` weight-only quantization modes. +- OpenAI-compatible unauthenticated HTTP API. +- Efficient generation by transferring activations during inference, not weights. + +## 2. Non-goals for the Initial Prototype + +The first prototype will intentionally avoid several advanced features: + +- No multi-model serving. +- No authentication. +- No request batching. +- No tensor parallelism across machines. +- No expert parallelism. +- No dynamic model hot-swapping. +- No high-performance CUDA-only quantization kernels as a requirement. +- No dependency on every node having HuggingFace access. +- No web UI. +- No streaming responses in the first milestone. +- No fault-tolerant recovery during an active generation. + +These can be added after a correct baseline works. + +## 3. High-Level Architecture + +TrueCluster uses pipeline-parallel inference. + +The coordinator owns: + +- CLI entry point for `cluster`. +- HTTP API. +- Tokenizer. +- Sampling logic. +- HuggingFace model metadata/config loading. +- Model weight loading from safetensors. +- Model split planning. +- Embedding layer. +- Final normalization. +- LM head. +- Node registry and orchestration. + +Each node owns: + +- CLI entry point for `node`. +- Persistent socket connection to coordinator. +- Device detection and selection. +- One contiguous range of transformer layers. +- KV cache for its assigned layers. +- Local forward execution on CUDA, MPS, or CPU. + +Generation path: + +```text +HTTP request + -> coordinator tokenizes prompt + -> coordinator runs embedding + -> hidden states sent to node 1 + -> node 1 runs assigned layers + -> hidden states sent to node 2 + -> ... + -> final hidden states returned to coordinator + -> coordinator runs final norm + lm_head + -> coordinator samples next token + -> repeat until completion +``` + +Weights are transferred once during assignment. During generation, only hidden states and small metadata are passed between coordinator and nodes. + +## 4. Execution Model + +### 4.1 Pipeline Parallelism + +The model is split by complete transformer blocks. Each node receives a contiguous set of layers: + +```text +coordinator: + embed_tokens + final_norm + lm_head + +node 1: + layers 0-7 + +node 2: + layers 8-15 + +node 3: + layers 16-23 +``` + +This is simpler and more reliable than tensor parallelism for heterogeneous machines and normal Ethernet/Wi-Fi networks. + +### 4.2 KV Cache Ownership + +KV cache is stored on the node that owns the relevant layers. + +For example: + +```text +node 1 cache: layers 0-7 +node 2 cache: layers 8-15 +node 3 cache: layers 16-23 +``` + +During prefill, each node creates cache entries for its layers. During decode, each node appends one token of keys/values to its cache. + +### 4.3 Request Concurrency + +Initial prototype supports one active generation at a time per cluster. + +Reason: distributed KV cache management is much easier with a single active request. Later versions can introduce request IDs, cache slots, and batching. + +## 5. CLI Design + +Use `typer` for the CLI. + +Package command: + +```bash +truecluster +``` + +### 5.1 Cluster Command + +Example: + +```bash +truecluster cluster \ + --model Qwen/Qwen3.5-0.8B \ + --node-host 0.0.0.0 \ + --node-port 7001 \ + --api-host 0.0.0.0 \ + --api-port 8000 \ + --max-nodes 4 \ + --quant fp16 +``` + +Options: + +```text +--model TEXT HuggingFace model id or local path. +--node-host TEXT Host/IP for worker-node socket server. Default: 0.0.0.0 +--node-port INT Port for worker-node socket server. Default: 7001 +--api-host TEXT Host/IP for HTTP API. Default: 0.0.0.0 +--api-port INT HTTP API port. Default: 8000 +--max-nodes INT Maximum number of worker nodes to use. +--quant [fp16|int8|int4] Weight mode. Default: fp16 +--dtype [fp16|bf16|fp32] Compute dtype preference. Default: fp16 +--target-node-memory-gb FLOAT Optional planning hint if node memory is unknown. +--trust-remote-code BOOL HuggingFace trust_remote_code. Default: false +--hf-cache-dir PATH Optional HuggingFace cache directory. +--log-level TEXT Default: info +``` + +Behavior: + +1. Resolve/download model. +2. Load config and tokenizer. +3. Load safetensors into coordinator RAM or memory-mapped index. +4. Build model tensor index. +5. Start node socket server. +6. Start HTTP API. +7. Wait for enough nodes. +8. Assign layer ranges and transmit weights. +9. Mark cluster ready. + +### 5.2 Node Command + +Example: + +```bash +truecluster node \ + --cluster-host 192.168.1.50 \ + --cluster-port 7001 \ + --device auto +``` + +Options: + +```text +--cluster-host TEXT Coordinator node socket host. +--cluster-port INT Coordinator node socket port. +--device TEXT auto, cuda, cuda:0, mps, or cpu. Default: auto +--node-id TEXT Optional stable node id. +--work-dir PATH Temporary local directory for received weights/cache. +--max-memory-gb FLOAT Optional memory capability override. +--log-level TEXT Default: info +``` + +Behavior: + +1. Detect hardware and PyTorch backends. +2. Connect to coordinator socket. +3. Send `HELLO` capability message. +4. Wait for assignment. +5. Receive config and layer weights. +6. Build local model shard. +7. Mark itself ready. +8. Serve prefill/decode requests over the persistent connection. + +## 6. HTTP API + +Use FastAPI and Uvicorn. + +The API is unauthenticated for the prototype. + +### 6.1 `GET /v1/models` + +Returns the single loaded model. + +Example response: + +```json +{ + "object": "list", + "data": [ + { + "id": "Qwen/Qwen3.5-0.8B", + "object": "model", + "created": 0, + "owned_by": "truecluster" + } + ] +} +``` + +### 6.2 `POST /v1/completions` + +Supported request fields initially: + +```json +{ + "model": "Qwen/Qwen3.5-0.8B", + "prompt": "Hello", + "max_tokens": 64, + "temperature": 0.7, + "top_p": 0.95, + "stop": null +} +``` + +Response should be OpenAI-compatible enough for common clients: + +```json +{ + "id": "cmpl-...", + "object": "text_completion", + "created": 0, + "model": "Qwen/Qwen3.5-0.8B", + "choices": [ + { + "text": " world", + "index": 0, + "logprobs": null, + "finish_reason": "stop" + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2 + } +} +``` + +### 6.3 `POST /v1/chat/completions` + +Supported request fields initially: + +```json +{ + "model": "Qwen/Qwen3.5-0.8B", + "messages": [ + {"role": "user", "content": "Hello"} + ], + "max_tokens": 64, + "temperature": 0.7, + "top_p": 0.95, + "stop": null, + "stream": false +} +``` + +The coordinator should use the HuggingFace tokenizer chat template if available. + +`stream: true` may return a clear unsupported error in the initial prototype. + +### 6.4 Not-Ready Error + +If a generation request arrives before enough nodes are connected and loaded: + +```json +{ + "error": { + "message": "Model is not ready. Required nodes: 3, connected ready nodes: 1", + "type": "cluster_not_ready" + } +} +``` + +HTTP status: `503`. + +## 7. Node Socket Protocol + +Use persistent TCP sockets with asyncio streams. + +Encoding: + +```text +[8-byte unsigned big-endian payload length][msgpack payload] +``` + +Large tensor payloads are sent as chunked binary data inside protocol messages or as msgpack metadata followed by raw bytes. + +### 7.1 Message Envelope + +Every message should include: + +```json +{ + "type": "MESSAGE_TYPE", + "request_id": "optional-request-id", + "seq": 1, + "payload": {} +} +``` + +### 7.2 Core Message Types + +Coordinator/node lifecycle: + +```text +HELLO +HELLO_ACK +ASSIGNMENT +WEIGHT_CHUNK +WEIGHTS_COMPLETE +LOAD_COMPLETE +LOAD_FAILED +PING +PONG +ERROR +``` + +Inference: + +```text +CLEAR_CACHE +RUN_PREFILL +RUN_DECODE +HIDDEN_STATE +INFERENCE_ERROR +``` + +### 7.3 `HELLO` + +Sent by node immediately after connection. + +Example: + +```json +{ + "type": "HELLO", + "payload": { + "node_id": "macbook-pro-1", + "hostname": "macbook-pro.local", + "python_version": "3.10.13", + "torch_version": "2.x", + "platform": "darwin", + "devices": [ + { + "id": "mps", + "type": "mps", + "name": "Apple Silicon MPS", + "total_memory": null, + "free_memory": null + } + ], + "selected_device": "mps", + "max_memory_bytes": null + } +} +``` + +CUDA example device: + +```json +{ + "id": "cuda:0", + "type": "cuda", + "name": "NVIDIA GeForce RTX 4090", + "total_memory": 25757220864, + "free_memory": 23000000000 +} +``` + +### 7.4 `ASSIGNMENT` + +Sent by coordinator after planning. + +```json +{ + "type": "ASSIGNMENT", + "payload": { + "model_id": "Qwen/Qwen3.5-0.8B", + "architecture": "qwen", + "quant": "fp16", + "compute_dtype": "fp16", + "layer_start": 0, + "layer_end_exclusive": 8, + "config": {}, + "tensor_count": 128, + "total_weight_bytes": 123456789 + } +} +``` + +### 7.5 Tensor Transfer + +Each tensor chunk message includes metadata: + +```json +{ + "type": "WEIGHT_CHUNK", + "payload": { + "tensor_name": "model.layers.0.self_attn.q_proj.weight", + "dtype": "float16", + "shape": [1024, 1024], + "chunk_index": 0, + "chunk_count": 4, + "offset": 0, + "data": "binary payload or external raw section" + } +} +``` + +For very large tensors, the preferred format is: + +```text +[frame length][msgpack metadata][raw bytes referenced by metadata] +``` + +The implementation should keep this hidden behind `protocol/tensors.py`. + +### 7.6 Inference Messages + +`RUN_PREFILL`: + +```json +{ + "type": "RUN_PREFILL", + "request_id": "req-1", + "payload": { + "position_start": 0, + "input_length": 42, + "hidden_state": "tensor payload" + } +} +``` + +`RUN_DECODE`: + +```json +{ + "type": "RUN_DECODE", + "request_id": "req-1", + "payload": { + "position": 42, + "hidden_state": "tensor payload" + } +} +``` + +`HIDDEN_STATE`: + +```json +{ + "type": "HIDDEN_STATE", + "request_id": "req-1", + "payload": { + "hidden_state": "tensor payload" + } +} +``` + +## 8. Model Support + +### 8.1 Initial Architecture + +Initial implementation should support Qwen2/Qwen2.5-style decoder-only causal LMs. The first implementation does not support Qwen3.5 hybrid models with `linear_attention`/GatedDeltaNet layers. + +Required components: + +- Token embedding. +- Stacked transformer decoder blocks. +- RMSNorm. +- Rotary position embeddings. +- Grouped-query attention. +- Causal attention mask. +- Gated MLP/SwiGLU. +- Final RMSNorm. +- LM head. +- Tied or untied output embeddings. + +### 8.2 HuggingFace Files + +Coordinator should support models with: + +```text +config.json +tokenizer.json/tokenizer.model/tokenizer_config.json +model.safetensors or model-00001-of-000xx.safetensors +model.safetensors.index.json, if sharded +``` + +Use libraries: + +- `huggingface_hub` +- `safetensors` +- `transformers` for tokenizer/config only where possible +- `torch` + +### 8.3 Weight Name Mapping + +For Qwen-style models, expected tensor names include patterns like: + +```text +model.embed_tokens.weight +model.layers.{i}.input_layernorm.weight +model.layers.{i}.self_attn.q_proj.weight +model.layers.{i}.self_attn.k_proj.weight +model.layers.{i}.self_attn.v_proj.weight +model.layers.{i}.self_attn.o_proj.weight +model.layers.{i}.post_attention_layernorm.weight +model.layers.{i}.mlp.gate_proj.weight +model.layers.{i}.mlp.up_proj.weight +model.layers.{i}.mlp.down_proj.weight +model.norm.weight +lm_head.weight +``` + +The model loader should validate required tensors before accepting a model. + +## 9. Model Planning and Splitting + +The coordinator computes layer assignments from model metadata and node capacity. + +### 9.1 Inputs + +- Number of transformer layers. +- Per-layer tensor byte sizes. +- Quantization mode. +- Max nodes. +- Connected node capabilities. +- Optional target memory hint. + +### 9.2 Rules + +- Split only on full transformer layer boundaries. +- Assign contiguous layer ranges. +- Preserve layer order. +- Coordinator keeps embeddings, final norm, and lm head. +- Required nodes must be connected and loaded before generation. +- If insufficient nodes are available, API returns `cluster_not_ready`. + +### 9.3 First Planner Algorithm + +Simple deterministic version: + +1. Compute total transformer layer bytes after quantization. +2. Estimate bytes per layer. +3. Determine number of partitions as `min(max_nodes, num_layers)`. +4. If node memory information is available, reduce/increase partitions so each assignment fits. +5. Otherwise split evenly by layer byte size. +6. Assign partitions to the first compatible ready nodes. + +### 9.4 Future Planner Improvements + +- Benchmark node speed and assign more layers to faster GPUs. +- Prefer CUDA nodes for larger shards. +- Consider network latency and bandwidth. +- Replicate small layers for resilience. +- Rebalance between generations. + +## 10. Quantization + +Quantization must work on CUDA, MPS, and CPU. Therefore the baseline implementation should avoid CUDA-only dependencies such as bitsandbytes. + +Supported modes: + +```text +fp16 +int8 +int4 +``` + +### 10.1 `fp16` + +- Store weights as `torch.float16`. +- Compute in `float16` by default. +- Works on CUDA and MPS. +- CPU fallback may use `float32` internally if needed. + +### 10.2 Portable `int8` + +Use symmetric per-output-channel weight-only quantization. + +For a linear weight `W` shaped `[out_features, in_features]`: + +```text +scale[out_features] = max(abs(W[row])) / 127 +qweight[row] = round(W[row] / scale[row]).clamp(-127, 127).int8 +``` + +Forward path: + +```text +W_dequant = qweight.float() * scale[:, None] +y = x @ W_dequant.T +``` + +This is portable but not maximally fast. + +### 10.3 Portable `int4` + +Use group-wise weight-only quantization. + +Suggested default group size: `128`. + +Store: + +```text +packed_qweight: uint8 +scale: float16/float32 per group +zero_point: optional +metadata: original shape, group size, packing order +``` + +Forward path: + +1. Unpack int4 values. +2. Dequantize to compute dtype. +3. Perform normal PyTorch matmul. + +This is designed for correctness and portability, not peak speed. + +### 10.4 Quantization Timing + +Preferred prototype behavior: + +- Coordinator loads original safetensors. +- Coordinator quantizes tensors before sending to nodes if `int8` or `int4` is selected. +- Nodes receive already-quantized tensors plus quantization metadata. +- Coordinator also quantizes/loads its own embedding/lm_head as needed. + +## 11. Generation Algorithm + +### 11.1 Prefill + +For prompt token IDs of length `N`: + +1. Coordinator computes embeddings: `[1, N, hidden_size]`. +2. Coordinator sends hidden state to first node with position start `0`. +3. Each node runs its assigned layers across the full sequence. +4. Each node initializes KV cache for its layers. +5. Final node returns hidden state to coordinator. +6. Coordinator applies final norm and lm head to the last token. +7. Coordinator samples next token. + +### 11.2 Decode + +For each generated token: + +1. Coordinator embeds last token: `[1, 1, hidden_size]`. +2. Coordinator sends hidden state to first node with current position. +3. Each node runs one-token decode using local KV cache. +4. Each node appends to its KV cache. +5. Final node returns hidden state to coordinator. +6. Coordinator computes logits and samples next token. +7. Stop if EOS, stop sequence, or `max_tokens` reached. + +### 11.3 Sampling + +Initial sampler supports: + +- Greedy when `temperature == 0`. +- Temperature scaling. +- Top-p nucleus sampling. +- EOS handling. +- Stop strings after decoding. + +Future additions: + +- Top-k. +- Repetition penalty. +- Frequency/presence penalties. +- Logprobs. + +## 12. Device Support + +### 12.1 Device Auto Detection + +Node device priority when `--device auto`: + +1. CUDA if available. +2. MPS if available. +3. CPU fallback. + +### 12.2 CUDA + +Use: + +```python +torch.cuda.is_available() +torch.cuda.get_device_properties(index) +torch.cuda.mem_get_info(index) +``` + +### 12.3 Apple MPS + +Use: + +```python +torch.backends.mps.is_available() +torch.device("mps") +``` + +MPS memory reporting is limited, so allow `--max-memory-gb` override. + +### 12.4 CPU + +CPU is allowed for testing and fallback, but may be slow. + +## 13. Package Structure + +Recommended source tree: + +```text +truecluster/ + __init__.py + cli.py + + cluster/ + __init__.py + server.py # node TCP server + api.py # FastAPI OpenAI-compatible API + planner.py # layer splitting + scheduler.py # generation orchestration + model_store.py # HF/safetensors loading + sampler.py + state.py + + node/ + __init__.py + client.py # connects to cluster + runtime.py # owns assigned layers + cache + device.py + + model/ + __init__.py + qwen.py # minimal Qwen implementation + layers.py + rotary.py + kv_cache.py + quant.py + tensor_names.py + + protocol/ + __init__.py + framing.py + messages.py + tensors.py + + tests/ + test_single_node_matches_transformers.py + test_protocol.py + test_quant.py + test_planner.py +``` + +Project metadata: + +```text +pyproject.toml +README.md +SPEC.md +``` + +## 14. Suggested Dependencies + +Runtime: + +```text +torch +transformers +huggingface_hub +safetensors +fastapi +uvicorn[standard] +typer +msgpack +pydantic +numpy +tqdm +``` + +Development/test: + +```text +pytest +pytest-asyncio +httpx +ruff +mypy optional +``` + +Python version: + +```text +>=3.10,<3.13 +``` + +## 15. Validation and Testing + +### 15.1 Correctness Test Against Transformers + +Most important validation: + +1. Load the target model with HuggingFace Transformers locally. +2. Load the same model through TrueCluster with one local node. +3. Run the same prompt. +4. Compare final logits before sampling. +5. Assert max difference is within tolerance for selected dtype. + +### 15.2 Multi-node Local Test + +Run on one machine: + +```bash +truecluster cluster --model ... --max-nodes 2 +truecluster node --cluster-host 127.0.0.1 --cluster-port 7001 --device cpu +truecluster node --cluster-host 127.0.0.1 --cluster-port 7001 --device cpu +``` + +Verify: + +- Both nodes receive different layer ranges. +- Prefill works. +- Decode works. +- Output matches single-node output within tolerance. + +### 15.3 Mixed Hardware Test + +Example: + +```text +coordinator: Mac or Linux host +node 1: Nvidia CUDA machine +node 2: Apple Silicon Mac MPS machine +``` + +Verify: + +- Both nodes connect. +- Assignments are sent. +- Generation completes. + +### 15.4 Quantization Tests + +For `int8` and `int4`: + +- Quantize/dequantize synthetic tensors. +- Check shape preservation. +- Check error bounds. +- Run short generation and verify no crashes. + +## 16. Implementation Phases + +### Phase 1: Local Single-Process Model Proof + +Deliverables: + +- Minimal Qwen model implementation. +- Safetensors loading. +- Local full-model forward. +- Logit comparison against Transformers. + +### Phase 2: One Node Distributed Inference + +Deliverables: + +- Socket protocol. +- Coordinator sends all transformer layers to one local node. +- Node loads layers and runs them. +- Coordinator keeps embedding/final norm/lm head. +- `/v1/completions` works. + +### Phase 3: Multi-node Layer Split + +Deliverables: + +- Planner assigns contiguous layer ranges. +- Multiple nodes are supported. +- Distributed KV cache works. +- Not-ready errors work. + +### Phase 4: Mixed CUDA/MPS Support + +Deliverables: + +- Device detection. +- CUDA execution. +- MPS execution. +- CPU fallback. +- Mixed Nvidia/Mac cluster generation test. + +### Phase 5: Portable Quantization + +Deliverables: + +- `fp16` baseline. +- Portable `int8` linear. +- Portable `int4` linear. +- Quantized weight transfer. +- CLI `--quant` option. + +### Phase 6: API Polish + +Deliverables: + +- `/v1/models`. +- `/v1/chat/completions`. +- Stop sequence support. +- Usage accounting. +- Better OpenAI-compatible errors. + +## 17. Initial Acceptance Criteria + +A prototype is considered working when: + +1. A cluster can be started with a HuggingFace safetensors Qwen-style model. +2. A node can connect to the cluster with no local model files. +3. The cluster sends layer weights to the node over the socket. +4. The node loads assigned layers on CUDA, MPS, or CPU. +5. `/v1/models` returns the loaded model. +6. `/v1/completions` generates text through the distributed pipeline. +7. `/v1/chat/completions` works for simple chat prompts. +8. If insufficient nodes are ready, API returns a clear `cluster_not_ready` error. +9. One local-node output matches Transformers logits within reasonable dtype tolerance. +10. Multi-node local CPU test works. + +## 18. Key Design Decisions + +- Use pipeline parallelism, not tensor parallelism. +- Use contiguous layer ranges only. +- Keep tokenizer, embeddings, final norm, lm head, and sampler on the coordinator. +- Send weights once at node assignment time. +- Send hidden states during generation. +- Store KV cache on worker nodes. +- Implement portable quantization instead of relying on CUDA-only libraries. +- Start with Qwen2/Qwen2.5-style causal LMs only. +- Optimize for correctness and clean architecture before speed. diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..80d76d8 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,36 @@ +[build-system] +requires = ["setuptools>=68", "wheel"] +build-backend = "setuptools.build_meta" + +[project] +name = "truecluster" +version = "0.1.0" +description = "Prototype heterogeneous pipeline-parallel LLM inference cluster" +readme = "README.md" +requires-python = ">=3.10,<3.13" +dependencies = [ + "torch", + "transformers", + "huggingface_hub", + "safetensors", + "fastapi", + "uvicorn[standard]", + "typer", + "msgpack", + "pydantic", + "numpy", + "tqdm", +] + +[project.optional-dependencies] +dev = ["pytest", "pytest-asyncio", "httpx", "ruff"] + +[project.scripts] +truecluster = "truecluster.cli:app" + +[tool.setuptools.packages.find] +include = ["truecluster*"] + +[tool.ruff] +line-length = 100 +target-version = "py310" diff --git a/tests/test_planner.py b/tests/test_planner.py new file mode 100644 index 0000000..3e55113 --- /dev/null +++ b/tests/test_planner.py @@ -0,0 +1,11 @@ +from truecluster.cluster.planner import plan_even_layers + + +def test_even_plan(): + assignments = plan_even_layers(num_layers=24, max_nodes=3) + assert [(a.layer_start, a.layer_end_exclusive) for a in assignments] == [(0, 8), (8, 16), (16, 24)] + + +def test_remainder_plan(): + assignments = plan_even_layers(num_layers=25, max_nodes=4) + assert [(a.layer_start, a.layer_end_exclusive) for a in assignments] == [(0, 7), (7, 13), (13, 19), (19, 25)] diff --git a/tests/test_tensors.py b/tests/test_tensors.py new file mode 100644 index 0000000..321e35e --- /dev/null +++ b/tests/test_tensors.py @@ -0,0 +1,11 @@ +import torch + +from truecluster.protocol.tensors import deserialize_tensor, serialize_tensor + + +def test_tensor_roundtrip_float16(): + x = torch.randn(2, 3).half() + y = deserialize_tensor(serialize_tensor(x)) + assert y.dtype == torch.float16 + assert y.shape == x.shape + assert torch.equal(x, y) diff --git a/truecluster/__init__.py b/truecluster/__init__.py new file mode 100644 index 0000000..3dc1f76 --- /dev/null +++ b/truecluster/__init__.py @@ -0,0 +1 @@ +__version__ = "0.1.0" diff --git a/truecluster/cli.py b/truecluster/cli.py new file mode 100644 index 0000000..50dfd98 --- /dev/null +++ b/truecluster/cli.py @@ -0,0 +1,107 @@ +from __future__ import annotations + +import asyncio +import logging +from typing import Optional + +import typer +import uvicorn + +from truecluster.cluster.api import create_app +from truecluster.cluster.model_store import ModelStore +from truecluster.cluster.server import ClusterRuntime, start_node_server +from truecluster.node.client import run_node + +app = typer.Typer(help="TrueCluster distributed LLM inference prototype") + + +def _setup_logging(level: str) -> None: + logging.basicConfig( + level=getattr(logging, level.upper(), logging.INFO), + format="%(asctime)s %(levelname)s [%(name)s] %(message)s", + ) + + +@app.command("cluster") +def cluster_cmd( + model: str = typer.Option(..., "--model", help="HuggingFace model id or local path"), + node_host: str = typer.Option("0.0.0.0", "--node-host", help="Worker socket host"), + node_port: int = typer.Option(7001, "--node-port", help="Worker socket port"), + api_host: str = typer.Option("0.0.0.0", "--api-host", help="HTTP API host"), + api_port: int = typer.Option(8000, "--api-port", help="HTTP API port"), + max_nodes: int = typer.Option(1, "--max-nodes", min=1, help="Maximum/requested worker shard count"), + quant: str = typer.Option("fp16", "--quant", help="Weight mode; currently only fp16 is implemented"), + target_node_memory_gb: Optional[float] = typer.Option(None, "--target-node-memory-gb", help="Optional planner memory hint"), + trust_remote_code: bool = typer.Option(False, "--trust-remote-code", help="Allow HF remote code for config/tokenizer"), + hf_cache_dir: Optional[str] = typer.Option(None, "--hf-cache-dir", help="Optional HuggingFace cache directory"), + log_level: str = typer.Option("info", "--log-level"), +) -> None: + _setup_logging(log_level) + asyncio.run( + _run_cluster( + model=model, + node_host=node_host, + node_port=node_port, + api_host=api_host, + api_port=api_port, + max_nodes=max_nodes, + quant=quant, + target_node_memory_gb=target_node_memory_gb, + trust_remote_code=trust_remote_code, + hf_cache_dir=hf_cache_dir, + ) + ) + + +async def _run_cluster( + model: str, + node_host: str, + node_port: int, + api_host: str, + api_port: int, + max_nodes: int, + quant: str, + target_node_memory_gb: float | None, + trust_remote_code: bool, + hf_cache_dir: str | None, +) -> None: + log = logging.getLogger("truecluster.cluster") + log.info("loading model into coordinator RAM: %s", model) + store = ModelStore.load(model, quant=quant, trust_remote_code=trust_remote_code, hf_cache_dir=hf_cache_dir) + log.info("model loaded: %s layers, %s tensors", store.num_layers, len(store.tensors)) + runtime = ClusterRuntime(store, max_nodes=max_nodes, target_node_memory_gb=target_node_memory_gb) + node_server = await start_node_server(node_host, node_port, runtime) + api = create_app(runtime) + config = uvicorn.Config(api, host=api_host, port=api_port, log_level="info") + http_server = uvicorn.Server(config) + log.info("HTTP API listening on %s:%s", api_host, api_port) + async with node_server: + await asyncio.gather(node_server.serve_forever(), http_server.serve()) + + +@app.command("node") +def node_cmd( + cluster_host: str = typer.Option(..., "--cluster-host", help="Coordinator socket host"), + cluster_port: int = typer.Option(7001, "--cluster-port", help="Coordinator socket port"), + device: str = typer.Option("auto", "--device", help="auto, cuda, cuda:0, mps, or cpu"), + node_id: Optional[str] = typer.Option(None, "--node-id", help="Optional stable node id"), + work_dir: Optional[str] = typer.Option(None, "--work-dir", help="Reserved for future temporary weight storage"), + max_memory_gb: Optional[float] = typer.Option(None, "--max-memory-gb", help="Optional capability override"), + log_level: str = typer.Option("info", "--log-level"), +) -> None: + _setup_logging(log_level) + if work_dir: + logging.getLogger(__name__).info("--work-dir is reserved for future use and is ignored in fp16 prototype") + asyncio.run( + run_node( + cluster_host=cluster_host, + cluster_port=cluster_port, + device=device, + node_id=node_id, + max_memory_gb=max_memory_gb, + ) + ) + + +if __name__ == "__main__": + app() diff --git a/truecluster/cluster/__init__.py b/truecluster/cluster/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/truecluster/cluster/api.py b/truecluster/cluster/api.py new file mode 100644 index 0000000..5b399cd --- /dev/null +++ b/truecluster/cluster/api.py @@ -0,0 +1,131 @@ +from __future__ import annotations + +import time +import uuid +from typing import Any + +from fastapi import FastAPI, HTTPException +from pydantic import BaseModel, Field + +from truecluster.cluster.server import ClusterRuntime + + +class CompletionRequest(BaseModel): + model: str | None = None + prompt: str | list[str] + max_tokens: int = Field(default=64, ge=1, le=4096) + temperature: float = 0.7 + top_p: float = 0.95 + stop: str | list[str] | None = None + stream: bool = False + + +class ChatMessage(BaseModel): + role: str + content: str + + +class ChatCompletionRequest(BaseModel): + model: str | None = None + messages: list[ChatMessage] + max_tokens: int = Field(default=64, ge=1, le=4096) + temperature: float = 0.7 + top_p: float = 0.95 + stop: str | list[str] | None = None + stream: bool = False + + +def create_app(runtime: ClusterRuntime) -> FastAPI: + app = FastAPI(title="TrueCluster", version="0.1.0") + + @app.get("/v1/models") + async def models() -> dict[str, Any]: + return { + "object": "list", + "data": [ + { + "id": runtime.model_store.model_id, + "object": "model", + "created": 0, + "owned_by": "truecluster", + } + ], + } + + @app.post("/v1/completions") + async def completions(req: CompletionRequest) -> dict[str, Any]: + if req.stream: + raise HTTPException(status_code=400, detail={"error": {"message": "streaming is not implemented", "type": "unsupported"}}) + prompt = req.prompt[0] if isinstance(req.prompt, list) else req.prompt + try: + result = await runtime.generate_completion( + prompt=prompt, + max_tokens=req.max_tokens, + temperature=req.temperature, + top_p=req.top_p, + stop=req.stop, + ) + except RuntimeError as exc: + raise HTTPException(status_code=503, detail={"error": {"message": str(exc), "type": "cluster_not_ready"}}) from exc + created = int(time.time()) + return { + "id": f"cmpl-{uuid.uuid4().hex}", + "object": "text_completion", + "created": created, + "model": runtime.model_store.model_id, + "choices": [ + { + "text": result["text"], + "index": 0, + "logprobs": None, + "finish_reason": result["finish_reason"], + } + ], + "usage": { + "prompt_tokens": result["prompt_tokens"], + "completion_tokens": result["completion_tokens"], + "total_tokens": result["total_tokens"], + }, + } + + @app.post("/v1/chat/completions") + async def chat_completions(req: ChatCompletionRequest) -> dict[str, Any]: + if req.stream: + raise HTTPException(status_code=400, detail={"error": {"message": "streaming is not implemented", "type": "unsupported"}}) + tokenizer = runtime.model_store.tokenizer + messages = [m.dict() for m in req.messages] + if hasattr(tokenizer, "apply_chat_template") and tokenizer.chat_template: + prompt = tokenizer.apply_chat_template(messages, tokenize=False, add_generation_prompt=True) + else: + prompt = "\n".join(f"{m['role']}: {m['content']}" for m in messages) + "\nassistant:" + try: + result = await runtime.generate_completion( + prompt=prompt, + max_tokens=req.max_tokens, + temperature=req.temperature, + top_p=req.top_p, + stop=req.stop, + ) + except RuntimeError as exc: + raise HTTPException(status_code=503, detail={"error": {"message": str(exc), "type": "cluster_not_ready"}}) from exc + created = int(time.time()) + return { + "id": f"chatcmpl-{uuid.uuid4().hex}", + "object": "chat.completion", + "created": created, + "model": runtime.model_store.model_id, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": result["text"]}, + "finish_reason": result["finish_reason"], + } + ], + "usage": { + "prompt_tokens": result["prompt_tokens"], + "completion_tokens": result["completion_tokens"], + "total_tokens": result["total_tokens"], + }, + } + + return app diff --git a/truecluster/cluster/model_store.py b/truecluster/cluster/model_store.py new file mode 100644 index 0000000..73d3a10 --- /dev/null +++ b/truecluster/cluster/model_store.py @@ -0,0 +1,178 @@ +from __future__ import annotations + +import json +from pathlib import Path +from typing import Any + +import torch +from huggingface_hub import snapshot_download +from safetensors.torch import load_file +from transformers import AutoConfig, AutoTokenizer + +from truecluster.model.qwen import CoordinatorHead, QwenConfig + + +class ModelStoreError(RuntimeError): + pass + + +class ModelStore: + """Coordinator-side model owner. + + The coordinator resolves the HF/local model, loads safetensors into RAM, and + keeps that RAM copy available for fast assignment transfer to worker nodes. + """ + + def __init__( + self, + model_id: str, + model_path: Path, + config_dict: dict[str, Any], + tokenizer: Any, + tensors: dict[str, torch.Tensor], + quant: str = "fp16", + ): + self.model_id = model_id + self.model_path = model_path + self.config_dict = config_dict + self.tokenizer = tokenizer + self.tensors = tensors + self.quant = quant + self.qwen_config = QwenConfig.from_dict(config_dict) + self.head = CoordinatorHead(config_dict, tensors, dtype=torch.float32) + + @classmethod + def load( + cls, + model: str, + quant: str = "fp16", + trust_remote_code: bool = False, + hf_cache_dir: str | None = None, + ) -> "ModelStore": + if quant != "fp16": + raise ModelStoreError("only fp16 is implemented currently; int8/int4 are planned") + + path = Path(model).expanduser() + if path.exists(): + model_path = path.resolve() + model_id = model + else: + model_path = Path( + snapshot_download( + repo_id=model, + cache_dir=hf_cache_dir, + allow_patterns=[ + "*.json", + "*.safetensors", + "*.model", + "tokenizer*", + "vocab*", + "merges.txt", + ], + ) + ) + model_id = model + + config = AutoConfig.from_pretrained(str(model_path), trust_remote_code=trust_remote_code) + config_dict = _normalize_supported_config(config.to_dict()) + tokenizer = AutoTokenizer.from_pretrained(str(model_path), trust_remote_code=trust_remote_code) + tensors = _load_safetensors_into_ram(model_path) + _convert_float_tensors_to_fp16(tensors) + _validate_qwen_tensors(config_dict, tensors) + return cls(model_id, model_path, config_dict, tokenizer, tensors, quant=quant) + + @property + def num_layers(self) -> int: + return self.qwen_config.num_hidden_layers + + def tensors_for_layers(self, layer_start: int, layer_end_exclusive: int) -> dict[str, torch.Tensor]: + out: dict[str, torch.Tensor] = {} + prefixes = [f"model.layers.{i}." for i in range(layer_start, layer_end_exclusive)] + for name, tensor in self.tensors.items(): + if any(name.startswith(prefix) for prefix in prefixes): + out[name] = tensor + return out + + def layer_bytes(self, layer_idx: int) -> int: + prefix = f"model.layers.{layer_idx}." + return sum(t.numel() * t.element_size() for name, t in self.tensors.items() if name.startswith(prefix)) + + +def _normalize_supported_config(config_dict: dict[str, Any]) -> dict[str, Any]: + """Return the text decoder config for architectures this prototype supports. + + The current runtime implements Qwen2/Qwen2.5-style full self-attention + decoder blocks. It intentionally does not implement Qwen3.5's hybrid + GatedDeltaNet/linear-attention blocks yet. + """ + + model_type = config_dict.get("model_type") + if model_type == "qwen3_5" or "text_config" in config_dict: + text_cfg = config_dict.get("text_config") or {} + layer_types = text_cfg.get("layer_types") or [] + if any(layer_type != "full_attention" for layer_type in layer_types): + raise ModelStoreError( + "Qwen3.5 hybrid/linear-attention models are not supported by the fp16 prototype yet. " + "Use a Qwen2/Qwen2.5 causal LM for now, for example " + "Qwen/Qwen2.5-0.5B-Instruct or Qwen/Qwen2.5-1.5B-Instruct." + ) + if text_cfg: + return text_cfg + + supported = {"qwen2", "qwen2_moe", "qwen3"} + if model_type not in supported: + raise ModelStoreError( + f"unsupported model_type {model_type!r}; current prototype supports Qwen2/Qwen2.5-style causal LMs" + ) + return config_dict + + +def _load_safetensors_into_ram(model_path: Path) -> dict[str, torch.Tensor]: + files = sorted(model_path.glob("*.safetensors")) + if not files: + raise ModelStoreError(f"no safetensors files found in {model_path}") + + tensors: dict[str, torch.Tensor] = {} + for file in files: + part = load_file(str(file), device="cpu") + overlap = set(tensors).intersection(part) + if overlap: + raise ModelStoreError(f"duplicate tensor names in safetensors: {sorted(overlap)[:5]}") + tensors.update(part) + return tensors + + +def _convert_float_tensors_to_fp16(tensors: dict[str, torch.Tensor]) -> None: + for name, tensor in list(tensors.items()): + if tensor.is_floating_point() and tensor.dtype != torch.float16: + tensors[name] = tensor.to(torch.float16).contiguous() + else: + tensors[name] = tensor.contiguous() + + +def _validate_qwen_tensors(config_dict: dict[str, Any], tensors: dict[str, torch.Tensor]) -> None: + cfg = QwenConfig.from_dict(config_dict) + required = ["model.embed_tokens.weight", "model.norm.weight"] + for i in range(cfg.num_hidden_layers): + p = f"model.layers.{i}" + required.extend( + [ + f"{p}.input_layernorm.weight", + f"{p}.self_attn.q_proj.weight", + f"{p}.self_attn.k_proj.weight", + f"{p}.self_attn.v_proj.weight", + f"{p}.self_attn.o_proj.weight", + f"{p}.post_attention_layernorm.weight", + f"{p}.mlp.gate_proj.weight", + f"{p}.mlp.up_proj.weight", + f"{p}.mlp.down_proj.weight", + ] + ) + missing = [name for name in required if name not in tensors] + if missing: + raise ModelStoreError("model does not look like supported Qwen-style safetensors; missing: " + ", ".join(missing[:20])) + + +def read_config_json(model_path: Path) -> dict[str, Any]: + with (model_path / "config.json").open("r", encoding="utf-8") as f: + return json.load(f) diff --git a/truecluster/cluster/planner.py b/truecluster/cluster/planner.py new file mode 100644 index 0000000..a19de5b --- /dev/null +++ b/truecluster/cluster/planner.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +from dataclasses import dataclass + + +@dataclass(frozen=True) +class LayerAssignment: + index: int + layer_start: int + layer_end_exclusive: int + + @property + def layer_count(self) -> int: + return self.layer_end_exclusive - self.layer_start + + +class PlannerError(RuntimeError): + pass + + +def plan_even_layers(num_layers: int, max_nodes: int, target_node_memory_bytes: int | None = None, layer_bytes: list[int] | None = None) -> list[LayerAssignment]: + if num_layers <= 0: + raise PlannerError("num_layers must be positive") + if max_nodes <= 0: + raise PlannerError("max_nodes must be positive") + + if target_node_memory_bytes and layer_bytes: + assignments: list[LayerAssignment] = [] + start = 0 + idx = 0 + while start < num_layers: + total = 0 + end = start + while end < num_layers and (total == 0 or total + layer_bytes[end] <= target_node_memory_bytes): + total += layer_bytes[end] + end += 1 + assignments.append(LayerAssignment(idx, start, end)) + idx += 1 + start = end + if len(assignments) > max_nodes: + raise PlannerError( + f"model needs {len(assignments)} nodes for target memory, but --max-nodes is {max_nodes}" + ) + return assignments + + # Without reliable node memory, use --max-nodes as the requested shard count. + partitions = min(max_nodes, num_layers) + base = num_layers // partitions + rem = num_layers % partitions + assignments = [] + start = 0 + for idx in range(partitions): + count = base + (1 if idx < rem else 0) + end = start + count + assignments.append(LayerAssignment(idx, start, end)) + start = end + return assignments diff --git a/truecluster/cluster/sampler.py b/truecluster/cluster/sampler.py new file mode 100644 index 0000000..66804cc --- /dev/null +++ b/truecluster/cluster/sampler.py @@ -0,0 +1,23 @@ +from __future__ import annotations + +import torch + + +@torch.inference_mode() +def sample_next_token(logits: torch.Tensor, temperature: float = 0.7, top_p: float = 0.95) -> int: + logits = logits[0, -1, :].float() + if temperature is None or temperature <= 0: + return int(torch.argmax(logits).item()) + logits = logits / float(temperature) + probs = torch.softmax(logits, dim=-1) + if top_p is not None and 0 < top_p < 1: + sorted_probs, sorted_indices = torch.sort(probs, descending=True) + cumulative = torch.cumsum(sorted_probs, dim=-1) + mask = cumulative > top_p + mask[1:] = mask[:-1].clone() + mask[0] = False + sorted_probs = sorted_probs.masked_fill(mask, 0.0) + sorted_probs = sorted_probs / sorted_probs.sum() + idx = torch.multinomial(sorted_probs, num_samples=1) + return int(sorted_indices[idx].item()) + return int(torch.multinomial(probs, num_samples=1).item()) diff --git a/truecluster/cluster/server.py b/truecluster/cluster/server.py new file mode 100644 index 0000000..c9fb040 --- /dev/null +++ b/truecluster/cluster/server.py @@ -0,0 +1,248 @@ +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: + 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 node.send(M.WEIGHT_TENSOR, {"name": name, "tensor": serialize_tensor(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 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: + log.exception("node connection failed from %s", peer) + 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 diff --git a/truecluster/model/__init__.py b/truecluster/model/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/truecluster/model/qwen.py b/truecluster/model/qwen.py new file mode 100644 index 0000000..d24f19b --- /dev/null +++ b/truecluster/model/qwen.py @@ -0,0 +1,252 @@ +from __future__ import annotations + +import math +from dataclasses import dataclass +from typing import Any + +import torch +import torch.nn.functional as F +from torch import nn + + +@dataclass +class QwenConfig: + vocab_size: int + hidden_size: int + intermediate_size: int + num_hidden_layers: int + num_attention_heads: int + num_key_value_heads: int + rms_norm_eps: float = 1e-6 + rope_theta: float = 1000000.0 + tie_word_embeddings: bool = False + max_position_embeddings: int = 32768 + + @property + def head_dim(self) -> int: + return self.hidden_size // self.num_attention_heads + + @classmethod + def from_dict(cls, data: dict[str, Any]) -> "QwenConfig": + return cls( + vocab_size=int(data["vocab_size"]), + hidden_size=int(data["hidden_size"]), + intermediate_size=int(data["intermediate_size"]), + num_hidden_layers=int(data["num_hidden_layers"]), + num_attention_heads=int(data["num_attention_heads"]), + num_key_value_heads=int(data.get("num_key_value_heads", data["num_attention_heads"])), + rms_norm_eps=float(data.get("rms_norm_eps", data.get("layer_norm_epsilon", 1e-6))), + rope_theta=float(data.get("rope_theta", 1000000.0)), + tie_word_embeddings=bool(data.get("tie_word_embeddings", False)), + max_position_embeddings=int(data.get("max_position_embeddings", 32768)), + ) + + +class RMSNorm(nn.Module): + def __init__(self, weight: torch.Tensor, eps: float): + super().__init__() + self.eps = eps + self.weight = nn.Parameter(weight, requires_grad=False) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + in_dtype = x.dtype + y = x.float() + y = y * torch.rsqrt(y.pow(2).mean(dim=-1, keepdim=True) + self.eps) + return (y.to(in_dtype) * self.weight.to(in_dtype)) + + +class LinearWeight(nn.Module): + def __init__(self, weight: torch.Tensor, bias: torch.Tensor | None = None): + super().__init__() + self.weight = nn.Parameter(weight, requires_grad=False) + self.bias = nn.Parameter(bias, requires_grad=False) if bias is not None else None + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return F.linear(x, self.weight.to(x.dtype), None if self.bias is None else self.bias.to(x.dtype)) + + +def _rotate_half(x: torch.Tensor) -> torch.Tensor: + x1 = x[..., : x.shape[-1] // 2] + x2 = x[..., x.shape[-1] // 2 :] + return torch.cat((-x2, x1), dim=-1) + + +def _rope_cache( + positions: torch.Tensor, + head_dim: int, + theta: float, + device: torch.device, + dtype: torch.dtype, +) -> tuple[torch.Tensor, torch.Tensor]: + inv_freq = 1.0 / (theta ** (torch.arange(0, head_dim, 2, device=device, dtype=torch.float32) / head_dim)) + freqs = torch.outer(positions.to(device=device, dtype=torch.float32), inv_freq) + emb = torch.cat((freqs, freqs), dim=-1) + return emb.cos().to(dtype), emb.sin().to(dtype) + + +def _apply_rope(x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: + # x: [B, H, T, D], cos/sin: [T, D] + cos = cos[None, None, :, :] + sin = sin[None, None, :, :] + return (x * cos) + (_rotate_half(x) * sin) + + +def _repeat_kv(x: torch.Tensor, repeats: int) -> torch.Tensor: + if repeats == 1: + return x + bsz, kv_heads, seq_len, head_dim = x.shape + x = x[:, :, None, :, :].expand(bsz, kv_heads, repeats, seq_len, head_dim) + return x.reshape(bsz, kv_heads * repeats, seq_len, head_dim) + + +class QwenAttention(nn.Module): + def __init__(self, config: QwenConfig, tensors: dict[str, torch.Tensor], prefix: str): + super().__init__() + self.config = config + self.num_heads = config.num_attention_heads + self.num_kv_heads = config.num_key_value_heads + self.head_dim = config.head_dim + self.num_kv_groups = self.num_heads // self.num_kv_heads + self.q_proj = LinearWeight(tensors[f"{prefix}.q_proj.weight"], tensors.get(f"{prefix}.q_proj.bias")) + self.k_proj = LinearWeight(tensors[f"{prefix}.k_proj.weight"], tensors.get(f"{prefix}.k_proj.bias")) + self.v_proj = LinearWeight(tensors[f"{prefix}.v_proj.weight"], tensors.get(f"{prefix}.v_proj.bias")) + self.o_proj = LinearWeight(tensors[f"{prefix}.o_proj.weight"], tensors.get(f"{prefix}.o_proj.bias")) + self.q_norm = RMSNorm(tensors[f"{prefix}.q_norm.weight"], config.rms_norm_eps) if f"{prefix}.q_norm.weight" in tensors else None + self.k_norm = RMSNorm(tensors[f"{prefix}.k_norm.weight"], config.rms_norm_eps) if f"{prefix}.k_norm.weight" in tensors else None + + def forward( + self, + x: torch.Tensor, + layer_cache: dict[str, torch.Tensor] | None, + positions: torch.Tensor, + ) -> tuple[torch.Tensor, dict[str, torch.Tensor]]: + bsz, seq_len, _ = x.shape + q = self.q_proj(x).view(bsz, seq_len, self.num_heads, self.head_dim).transpose(1, 2) + k = self.k_proj(x).view(bsz, seq_len, self.num_kv_heads, self.head_dim).transpose(1, 2) + v = self.v_proj(x).view(bsz, seq_len, self.num_kv_heads, self.head_dim).transpose(1, 2) + + if self.q_norm is not None: + q = self.q_norm(q) + if self.k_norm is not None: + k = self.k_norm(k) + + cos, sin = _rope_cache(positions, self.head_dim, self.config.rope_theta, x.device, q.dtype) + q = _apply_rope(q, cos, sin) + k = _apply_rope(k, cos, sin) + + if layer_cache is not None and "k" in layer_cache: + k_all = torch.cat([layer_cache["k"], k], dim=2) + v_all = torch.cat([layer_cache["v"], v], dim=2) + else: + k_all = k + v_all = v + new_cache = {"k": k_all.detach(), "v": v_all.detach()} + + k_rep = _repeat_kv(k_all, self.num_kv_groups) + v_rep = _repeat_kv(v_all, self.num_kv_groups) + + scores = torch.matmul(q.float(), k_rep.float().transpose(-2, -1)) / math.sqrt(self.head_dim) + total_len = k_rep.shape[-2] + key_positions = torch.arange(total_len, device=x.device)[None, None, None, :] + query_positions = positions.to(x.device)[None, None, :, None] + scores = scores.masked_fill(key_positions > query_positions, torch.finfo(scores.dtype).min) + attn = torch.softmax(scores, dim=-1).to(q.dtype) + out = torch.matmul(attn, v_rep) + out = out.transpose(1, 2).contiguous().view(bsz, seq_len, self.config.hidden_size) + return self.o_proj(out), new_cache + + +class QwenMLP(nn.Module): + def __init__(self, tensors: dict[str, torch.Tensor], prefix: str): + super().__init__() + self.gate_proj = LinearWeight(tensors[f"{prefix}.gate_proj.weight"], tensors.get(f"{prefix}.gate_proj.bias")) + self.up_proj = LinearWeight(tensors[f"{prefix}.up_proj.weight"], tensors.get(f"{prefix}.up_proj.bias")) + self.down_proj = LinearWeight(tensors[f"{prefix}.down_proj.weight"], tensors.get(f"{prefix}.down_proj.bias")) + + def forward(self, x: torch.Tensor) -> torch.Tensor: + return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) + + +class QwenDecoderLayer(nn.Module): + def __init__(self, config: QwenConfig, tensors: dict[str, torch.Tensor], layer_idx: int): + super().__init__() + prefix = f"model.layers.{layer_idx}" + self.layer_idx = layer_idx + self.input_layernorm = RMSNorm(tensors[f"{prefix}.input_layernorm.weight"], config.rms_norm_eps) + self.self_attn = QwenAttention(config, tensors, f"{prefix}.self_attn") + self.post_attention_layernorm = RMSNorm(tensors[f"{prefix}.post_attention_layernorm.weight"], config.rms_norm_eps) + self.mlp = QwenMLP(tensors, f"{prefix}.mlp") + + def forward( + self, + x: torch.Tensor, + cache: dict[int, dict[str, torch.Tensor]], + positions: torch.Tensor, + ) -> torch.Tensor: + residual = x + attn_out, new_cache = self.self_attn(self.input_layernorm(x), cache.get(self.layer_idx), positions) + cache[self.layer_idx] = new_cache + x = residual + attn_out + residual = x + x = residual + self.mlp(self.post_attention_layernorm(x)) + return x + + +class QwenLayerShard(nn.Module): + def __init__( + self, + config_dict: dict[str, Any], + layer_start: int, + layer_end_exclusive: int, + tensors: dict[str, torch.Tensor], + device: str | torch.device, + dtype: torch.dtype = torch.float16, + ): + super().__init__() + self.config = QwenConfig.from_dict(config_dict) + self.layer_start = layer_start + self.layer_end_exclusive = layer_end_exclusive + self.device = torch.device(device) + self.dtype = dtype + local_tensors = {k: v.to(self.device, dtype=dtype if v.is_floating_point() else v.dtype) for k, v in tensors.items()} + self.layers = nn.ModuleList( + [QwenDecoderLayer(self.config, local_tensors, i) for i in range(layer_start, layer_end_exclusive)] + ) + self.cache: dict[int, dict[str, torch.Tensor]] = {} + self.to(self.device) + self.eval() + + def clear_cache(self) -> None: + self.cache.clear() + + @torch.inference_mode() + def forward(self, hidden: torch.Tensor, position_start: int) -> torch.Tensor: + hidden = hidden.to(self.device, dtype=self.dtype) + seq_len = hidden.shape[1] + positions = torch.arange(position_start, position_start + seq_len, device=self.device, dtype=torch.long) + for layer in self.layers: + hidden = layer(hidden, self.cache, positions) + return hidden + + +class CoordinatorHead: + def __init__(self, config_dict: dict[str, Any], tensors: dict[str, torch.Tensor], dtype: torch.dtype = torch.float32): + self.config = QwenConfig.from_dict(config_dict) + self.dtype = dtype + self.embed_weight = tensors["model.embed_tokens.weight"].to("cpu", dtype=dtype) + self.norm_weight = tensors["model.norm.weight"].to("cpu", dtype=dtype) + self.lm_head_weight = tensors.get("lm_head.weight", tensors["model.embed_tokens.weight"]).to("cpu", dtype=dtype) + self.eps = self.config.rms_norm_eps + + @torch.inference_mode() + def embed(self, input_ids: torch.Tensor) -> torch.Tensor: + input_ids = input_ids.to("cpu", dtype=torch.long) + return F.embedding(input_ids, self.embed_weight) + + @torch.inference_mode() + def logits(self, hidden: torch.Tensor) -> torch.Tensor: + hidden = hidden.to("cpu", dtype=self.dtype) + y = hidden.float() + y = y * torch.rsqrt(y.pow(2).mean(dim=-1, keepdim=True) + self.eps) + y = y.to(self.dtype) * self.norm_weight + return F.linear(y, self.lm_head_weight) diff --git a/truecluster/node/__init__.py b/truecluster/node/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/truecluster/node/client.py b/truecluster/node/client.py new file mode 100644 index 0000000..6257b8c --- /dev/null +++ b/truecluster/node/client.py @@ -0,0 +1,79 @@ +from __future__ import annotations + +import asyncio +import logging + +from truecluster.node.device import capabilities, select_device +from truecluster.node.runtime import NodeRuntime +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__) + + +async def run_node( + cluster_host: str, + cluster_port: int, + device: str = "auto", + node_id: str | None = None, + max_memory_gb: float | None = None, +) -> None: + selected = select_device(device) + runtime = NodeRuntime(selected) + log.info("connecting to cluster %s:%s using device %s", cluster_host, cluster_port, selected) + reader, writer = await asyncio.open_connection(cluster_host, cluster_port) + await write_message(writer, make_message(M.HELLO, capabilities(selected, node_id=node_id, max_memory_gb=max_memory_gb))) + try: + while True: + msg = await read_message(reader) + if msg is None: + log.warning("cluster disconnected") + return + msg_type = msg.get("type") + payload = msg.get("payload", {}) + try: + if msg_type == M.HELLO_ACK: + log.info("connected to cluster; required_nodes=%s", payload.get("required_nodes")) + elif msg_type == M.ASSIGNMENT: + runtime.set_assignment(payload) + log.info( + "received assignment: layers [%s,%s), tensors=%s, bytes=%.2f MB", + payload.get("layer_start"), + payload.get("layer_end_exclusive"), + payload.get("tensor_count"), + int(payload.get("total_weight_bytes", 0)) / 1024 / 1024, + ) + elif msg_type == M.WEIGHT_TENSOR: + runtime.add_tensor(payload["name"], deserialize_tensor(payload["tensor"], device="cpu")) + elif msg_type == M.WEIGHTS_COMPLETE: + log.info("all weights received; loading shard") + runtime.load() + await write_message(writer, make_message(M.LOAD_COMPLETE, {"device": selected})) + log.info("shard loaded and ready") + elif msg_type == M.CLEAR_CACHE: + runtime.clear_cache() + elif msg_type in (M.RUN_PREFILL, M.RUN_DECODE): + hidden = deserialize_tensor(payload["hidden_state"], device="cpu") + position_start = int(payload["position_start"]) + out = runtime.forward(hidden, position_start=position_start) + await write_message( + writer, + make_message( + M.HIDDEN_STATE, + {"hidden_state": serialize_tensor(out)}, + request_id=msg.get("request_id"), + ), + ) + elif msg_type == M.PING: + await write_message(writer, make_message(M.PONG, {})) + elif msg_type == M.ERROR: + log.error("cluster error: %s", payload) + else: + log.warning("unknown message type from cluster: %s", msg_type) + except Exception as exc: + log.exception("failed handling message %s", msg_type) + await write_message(writer, make_message(M.LOAD_FAILED if msg_type in (M.WEIGHTS_COMPLETE, M.ASSIGNMENT) else M.ERROR, {"message": str(exc)})) + finally: + writer.close() + await writer.wait_closed() diff --git a/truecluster/node/device.py b/truecluster/node/device.py new file mode 100644 index 0000000..697f39f --- /dev/null +++ b/truecluster/node/device.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +import platform +import socket +import sys +from typing import Any + +import torch + + +def select_device(requested: str = "auto") -> str: + requested = requested.lower() + if requested == "auto": + if torch.cuda.is_available(): + return "cuda:0" + if getattr(torch.backends, "mps", None) is not None and torch.backends.mps.is_available(): + return "mps" + return "cpu" + if requested == "cuda": + return "cuda:0" + return requested + + +def capabilities(selected_device: str, node_id: str | None = None, max_memory_gb: float | None = None) -> dict[str, Any]: + devices: list[dict[str, Any]] = [] + if torch.cuda.is_available(): + for i in range(torch.cuda.device_count()): + props = torch.cuda.get_device_properties(i) + free = total = None + try: + free, total = torch.cuda.mem_get_info(i) + except Exception: + total = props.total_memory + devices.append( + { + "id": f"cuda:{i}", + "type": "cuda", + "name": props.name, + "total_memory": int(total) if total is not None else None, + "free_memory": int(free) if free is not None else None, + } + ) + if getattr(torch.backends, "mps", None) is not None and torch.backends.mps.is_available(): + devices.append( + { + "id": "mps", + "type": "mps", + "name": "Apple Silicon MPS", + "total_memory": None, + "free_memory": None, + } + ) + devices.append({"id": "cpu", "type": "cpu", "name": platform.processor() or "CPU", "total_memory": None, "free_memory": None}) + return { + "node_id": node_id or socket.gethostname(), + "hostname": socket.gethostname(), + "python_version": sys.version.split()[0], + "torch_version": torch.__version__, + "platform": sys.platform, + "devices": devices, + "selected_device": selected_device, + "max_memory_bytes": int(max_memory_gb * 1024**3) if max_memory_gb else None, + } diff --git a/truecluster/node/runtime.py b/truecluster/node/runtime.py new file mode 100644 index 0000000..44e8e2c --- /dev/null +++ b/truecluster/node/runtime.py @@ -0,0 +1,46 @@ +from __future__ import annotations + +from typing import Any + +import torch + +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.shard: QwenLayerShard | None = None + + def set_assignment(self, payload: dict[str, Any]) -> None: + self.assignment = payload + self.tensors = {} + self.shard = None + + def add_tensor(self, name: str, tensor: torch.Tensor) -> None: + self.tensors[name] = tensor + + 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() diff --git a/truecluster/protocol/__init__.py b/truecluster/protocol/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/truecluster/protocol/framing.py b/truecluster/protocol/framing.py new file mode 100644 index 0000000..72e70a1 --- /dev/null +++ b/truecluster/protocol/framing.py @@ -0,0 +1,52 @@ +from __future__ import annotations + +import asyncio +import struct +from typing import Any + +import msgpack + +MAX_FRAME_BYTES = 4 * 1024 * 1024 * 1024 # 4 GiB; practical frames should be much smaller. + + +class ProtocolError(RuntimeError): + pass + + +async def read_message(reader: asyncio.StreamReader) -> dict[str, Any] | None: + """Read one length-prefixed msgpack message. + + Returns None on clean EOF before the frame header. + """ + + try: + header = await reader.readexactly(8) + except asyncio.IncompleteReadError as exc: + if not exc.partial: + return None + raise ProtocolError("incomplete frame header") from exc + + (length,) = struct.unpack(">Q", header) + if length > MAX_FRAME_BYTES: + raise ProtocolError(f"frame too large: {length} bytes") + try: + data = await reader.readexactly(length) + except asyncio.IncompleteReadError as exc: + raise ProtocolError("incomplete frame payload") from exc + msg = msgpack.unpackb(data, raw=False, use_list=True, strict_map_key=False) + if not isinstance(msg, dict): + raise ProtocolError("message must be a map") + return msg + + +async def write_message(writer: asyncio.StreamWriter, message: dict[str, Any]) -> None: + data = msgpack.packb(message, use_bin_type=True) + writer.write(struct.pack(">Q", len(data)) + data) + await writer.drain() + + +def make_message(msg_type: str, payload: dict[str, Any] | None = None, request_id: str | None = None) -> dict[str, Any]: + msg: dict[str, Any] = {"type": msg_type, "payload": payload or {}} + if request_id is not None: + msg["request_id"] = request_id + return msg diff --git a/truecluster/protocol/messages.py b/truecluster/protocol/messages.py new file mode 100644 index 0000000..3fc7327 --- /dev/null +++ b/truecluster/protocol/messages.py @@ -0,0 +1,14 @@ +HELLO = "HELLO" +HELLO_ACK = "HELLO_ACK" +ASSIGNMENT = "ASSIGNMENT" +WEIGHT_TENSOR = "WEIGHT_TENSOR" +WEIGHTS_COMPLETE = "WEIGHTS_COMPLETE" +LOAD_COMPLETE = "LOAD_COMPLETE" +LOAD_FAILED = "LOAD_FAILED" +CLEAR_CACHE = "CLEAR_CACHE" +RUN_PREFILL = "RUN_PREFILL" +RUN_DECODE = "RUN_DECODE" +HIDDEN_STATE = "HIDDEN_STATE" +ERROR = "ERROR" +PING = "PING" +PONG = "PONG" diff --git a/truecluster/protocol/tensors.py b/truecluster/protocol/tensors.py new file mode 100644 index 0000000..5ff305e --- /dev/null +++ b/truecluster/protocol/tensors.py @@ -0,0 +1,88 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Any + +import numpy as np +import torch + +_DTYPE_TO_TORCH = { + "float16": torch.float16, + "float32": torch.float32, + "float64": torch.float64, + "int64": torch.int64, + "int32": torch.int32, + "int16": torch.int16, + "int8": torch.int8, + "uint8": torch.uint8, + "bool": torch.bool, +} + +_TORCH_TO_NUMPY = { + torch.float16: np.float16, + torch.float32: np.float32, + torch.float64: np.float64, + torch.int64: np.int64, + torch.int32: np.int32, + torch.int16: np.int16, + torch.int8: np.int8, + torch.uint8: np.uint8, + torch.bool: np.bool_, +} + + +@dataclass(frozen=True) +class SerializedTensor: + dtype: str + shape: list[int] + data: bytes + + def to_payload(self) -> dict[str, Any]: + return {"dtype": self.dtype, "shape": self.shape, "data": self.data} + + +class TensorSerializationError(RuntimeError): + pass + + +def dtype_name(dtype: torch.dtype) -> str: + text = str(dtype) + if text.startswith("torch."): + return text.split(".", 1)[1] + return text + + +def serialize_tensor(tensor: torch.Tensor) -> dict[str, Any]: + """Serialize a tensor into msgpack-friendly metadata + raw bytes. + + Tensors are moved to CPU and made contiguous. bfloat16 is converted to float16 + because NumPy has inconsistent bfloat16 support and the prototype currently + targets fp16 transfer. + """ + + tensor = tensor.detach().cpu().contiguous() + if tensor.dtype == torch.bfloat16: + tensor = tensor.to(torch.float16) + if tensor.dtype not in _TORCH_TO_NUMPY: + raise TensorSerializationError(f"unsupported tensor dtype for serialization: {tensor.dtype}") + array = tensor.numpy() + return { + "dtype": dtype_name(tensor.dtype), + "shape": list(tensor.shape), + "data": array.tobytes(order="C"), + } + + +def deserialize_tensor(payload: dict[str, Any], device: str | torch.device | None = None) -> torch.Tensor: + dtype = payload["dtype"] + shape = list(payload["shape"]) + data = payload["data"] + torch_dtype = _DTYPE_TO_TORCH.get(dtype) + if torch_dtype is None: + raise TensorSerializationError(f"unsupported tensor dtype for deserialization: {dtype}") + np_dtype = _TORCH_TO_NUMPY[torch_dtype] + array = np.frombuffer(data, dtype=np_dtype).copy().reshape(shape) + tensor = torch.from_numpy(array) + if device is not None: + tensor = tensor.to(device) + return tensor