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