Inital commit
This commit is contained in:
+47
@@ -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/
|
||||||
@@ -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}'
|
||||||
|
```
|
||||||
@@ -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.
|
||||||
@@ -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"
|
||||||
@@ -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)]
|
||||||
@@ -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)
|
||||||
@@ -0,0 +1 @@
|
|||||||
|
__version__ = "0.1.0"
|
||||||
@@ -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()
|
||||||
@@ -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
|
||||||
@@ -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)
|
||||||
@@ -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
|
||||||
@@ -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())
|
||||||
@@ -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
|
||||||
@@ -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)
|
||||||
@@ -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()
|
||||||
@@ -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,
|
||||||
|
}
|
||||||
@@ -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()
|
||||||
@@ -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
|
||||||
@@ -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"
|
||||||
@@ -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
|
||||||
Reference in New Issue
Block a user