Improve CUDA setup and chat request handling

This commit is contained in:
2026-06-05 12:56:30 -05:00
parent 04eafb5446
commit 3b2a3e85de
3 changed files with 83 additions and 2 deletions
+64
View File
@@ -10,6 +10,70 @@ See [`SPEC.md`](SPEC.md).
pip install -e .
```
### Nvidia CUDA install
If a node has an Nvidia GPU, it must install a CUDA-enabled PyTorch build. If you see:
```text
Torch not compiled with CUDA enabled
```
then the node installed the CPU-only PyTorch package.
Recommended fix on a macOS/Linux Nvidia node:
```bash
source .venv/bin/activate
pip uninstall -y torch torchvision torchaudio
pip install --index-url https://download.pytorch.org/whl/cu121 torch torchvision torchaudio
pip install -e .
```
Recommended fix on a Windows Nvidia node using PowerShell:
```powershell
.\.venv\Scripts\Activate.ps1
python -m pip uninstall -y torch torchvision torchaudio
python -m pip install --index-url https://download.pytorch.org/whl/cu121 torch torchvision torchaudio
python -m pip install -e .
```
If PowerShell blocks venv activation, run this once in the same PowerShell window:
```powershell
Set-ExecutionPolicy -Scope Process -ExecutionPolicy Bypass
.\.venv\Scripts\Activate.ps1
```
Windows Command Prompt alternative:
```bat
.venv\Scripts\activate.bat
python -m pip uninstall -y torch torchvision torchaudio
python -m pip install --index-url https://download.pytorch.org/whl/cu121 torch torchvision torchaudio
python -m pip install -e .
```
For newer CUDA builds, PyTorch may also provide `cu124` or `cu126` wheels. Check https://pytorch.org/get-started/locally/ if `cu121` is not appropriate.
Verify CUDA support:
```bash
python - <<'PY'
import torch
print('torch:', torch.__version__)
print('cuda available:', torch.cuda.is_available())
print('cuda version:', torch.version.cuda)
print('gpu:', torch.cuda.get_device_name(0) if torch.cuda.is_available() else None)
PY
```
Then run the node with:
```bash
truecluster node --cluster-host YOUR_CLUSTER_IP --cluster-port 7001 --device cuda:0
```
## Run a cluster
```bash
+17 -1
View File
@@ -52,8 +52,21 @@ def create_app(runtime: ClusterRuntime) -> FastAPI:
],
}
def _validate_model(requested: str | None) -> None:
if requested and requested != runtime.model_store.model_id:
raise HTTPException(
status_code=400,
detail={
"error": {
"message": f"requested model {requested!r} does not match loaded model {runtime.model_store.model_id!r}",
"type": "model_mismatch",
}
},
)
@app.post("/v1/completions")
async def completions(req: CompletionRequest) -> dict[str, Any]:
_validate_model(req.model)
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
@@ -90,11 +103,13 @@ def create_app(runtime: ClusterRuntime) -> FastAPI:
@app.post("/v1/chat/completions")
async def chat_completions(req: ChatCompletionRequest) -> dict[str, Any]:
_validate_model(req.model)
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:
used_chat_template = bool(hasattr(tokenizer, "apply_chat_template") and tokenizer.chat_template)
if used_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:"
@@ -105,6 +120,7 @@ def create_app(runtime: ClusterRuntime) -> FastAPI:
temperature=req.temperature,
top_p=req.top_p,
stop=req.stop,
add_special_tokens=not used_chat_template,
)
except RuntimeError as exc:
raise HTTPException(status_code=503, detail={"error": {"message": str(exc), "type": "cluster_not_ready"}}) from exc
+2 -1
View File
@@ -154,6 +154,7 @@ class ClusterRuntime:
temperature: float = 0.7,
top_p: float = 0.95,
stop: str | list[str] | None = None,
add_special_tokens: bool = True,
) -> dict[str, Any]:
from truecluster.cluster.sampler import sample_next_token
@@ -163,7 +164,7 @@ class ClusterRuntime:
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)
encoded = tokenizer(prompt, return_tensors="pt", add_special_tokens=add_special_tokens)
input_ids = encoded["input_ids"].to(torch.long)
prompt_tokens = int(input_ids.shape[1])