Files
rustytorch/CLAUDE.md
T
osobhandClaude Fable 5 fce6cef262
CI / Clippy Check (push) Failing after 5s
CI / Format Check (push) Failing after 5s
CI / Build CPU-Only (Explicit) (push) Failing after 6s
CI / Build (ubuntu-latest) (push) Failing after 6s
Documentation / Build User Guide (push) Successful in 4s
CI / Build (macos-latest) (push) Failing after 7s
CI / Test (macos-latest) (push) Has been skipped
CI / Test (ubuntu-latest) (push) Has been skipped
CI / Python Bindings (maturin) (macos-latest) (push) Has been skipped
CI / Python Bindings (maturin) (ubuntu-latest) (push) Has been skipped
CI / WASM Build + Size Check (push) Has been skipped
CI / Distributed Training Tests (push) Has been skipped
CI / CI Success (push) Failing after 0s
Performance Benchmarks / Run Benchmarks (push) Successful in 25s
Documentation / Build API Documentation (push) Failing after 29s
docs: JEPA roadmap — GPU resume re-upload done; remaining items need multi-GPU
Co-Authored-By: Claude Fable 5 <[email protected]>
2026-07-10 03:59:48 -07:00

164 lines
11 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# rustytorch
> GPU-accelerated ML framework in pure Rust — full PyTorch-equivalent with CUDA/Metal/ROCm/WebGPU backends, 26 batches of optimizations (19 Blackwell SM_120 GPU perf + 7 JEPA platform), complete transformer training arsenal (Muon/Shampoo/SOAP/ScheduleFree/DPO/TIES-DARE/MoD), vLLM-class inference (Medusa/EAGLE/Lookahead speculative decoding, paged KV cache, chunked prefill), federated learning, domain-specific stacks for medical imaging/neuroimaging/scientific computing, and a complete JEPA self-supervised learning platform (I-JEPA, V-JEPA, Neuro-JEPA, ViT wiring, data pipeline, cluster config — Batches 2026). Primary goal: premier JEPA platform for multi-node GPU cluster.
## Problems It Solves
- No production-ready Rust ML framework matches PyTorch's feature breadth with multi-backend GPU support — rustytorch fills this gap with memory safety and zero GIL limitations
- Flash Attention for 128K+ token context windows requires memory-efficient attention kernels not available in generic frameworks
- Speculative decoding (23× inference speedup) requires tight integration between draft and verification models — implemented as a first-class feature
- Domain-specific workloads (DICOM medical imaging, EEG/MEG neuroimaging, PDE-solving via Neural Operators) require specialized data pipelines and architectures beyond generic ML frameworks
- Distributed training across multi-GPU clusters requires NCCL-level collective communication integrated directly into the training loop
## What It's Comprised Of
113 crates organized across 8 layers (GPU Perf Batches 119 complete as of 2026-06-27):
### Layer 1: Core Infrastructure (26 crates)
| Crate | Role |
|-------|------|
| `rtx-backend` | Multi-backend abstraction (CUDA, Metal, ROCm, SYCL, WebGPU, CPU) |
| `rtx-tensor` | Zero-copy GPU-native tensor primitives, PyTorch API compatible |
| `rtx-autograd` | Tape-based automatic differentiation with Hessian-vector products |
| `rtx-memory` | Memory pooling, NUMA awareness, huge pages |
| `rtx-kernel` | GPU kernel compilation and execution (CUDA PTX + Metal MSL) |
| `rtx-flash-attention` | Memory-efficient Flash Attention for 128K+ token contexts |
| `rtx-lora` | Low-Rank Adapter (LoRA/QLoRA) fine-tuning |
| `rtx-moe` | Mixture of Experts (SwitchTransformer, ExpertChoice, TokenChoice routing) |
| `rtx-mamba` | Mamba/State Space Models (O(n) attention alternative, selective scan) |
| `rtx-inference::speculative` | Speculative decoding for 23× inference speedup (module in rtx-inference, not a separate crate) |
| `rtx-fusion` | Kernel fusion DSL eliminating intermediate tensor writes |
| `rtx-cubecl` | CubeCL integration for Rust → GPU kernel synthesis |
| `rtx-bindings` | PyO3 Python bindings, C FFI, ONNX/DLPack interop |
| `rtx-macros` | Derive macros (`Module`, `Config`) for ergonomic neural network API |
### Layer 2: Training & Optimization (16 crates)
| Crate | Role |
|-------|------|
| `rtx-transformers` | Complete transformer: MQA, GQA, SwiGLU, RoPE, ALiBi positional encodings |
| `rtx-distributed` | Multi-GPU/multi-node training with NCCL/RCCL, elastic recovery |
| `rtx-rl` | Reinforcement Learning: PPO, RLHF, reward modeling |
| `rtx-compress` | Model compression: pruning, distillation, AWQ/GPTQ/SmoothQuant quantization |
| `rtx-nas` | Neural Architecture Search: DARTS, PC-DARTS, FairNAS |
| `rtx-federated` | Federated learning for 100,000+ edge devices |
| `rtx-auto` | AutoML with hyperparameter optimization |
| `rtx-evolution` | Evolutionary optimization strategies |
### Layer 3: Model Architectures (10 crates)
| Crate | Role |
|-------|------|
| `rtx-vision` | Vision Transformers (ViT, ConvNeXt, patch embeddings) |
| `rtx-multimodal` | Vision-language models (CLIP-style, audio transformers) |
| `rtx-diffuse` | Diffusion models (DDIM, classifier-free guidance) |
| `rtx-timeseries` | Time series: ARIMA, Prophet, state-space models |
| `rtx-nlg` | Natural language generation, GPT-style language models |
### Layer 4: Production & Deployment (9 crates)
| Crate | Role |
|-------|------|
| `rtx-serving-api` | HTTP/gRPC model serving with load balancing |
| `rtx-inference` | High-performance inference: continuous batching, paged KV cache |
| `rtx-streaming` | Real-time ML inference with <1ms latency |
| `rtx-wasm-inference` | WebAssembly inference for browser deployment |
| `rtx-onnx` | ONNX model import/export |
### Layer 5: Specialized Domains (29 crates)
| Domain | Key Crates |
|--------|-----------|
| Medical imaging | `rtx-medical-core`, `rtx-mri2fe`, `rtx-registration`, `rtx-segmentation` (DICOM, NIfTI) |
| Neuroimaging | `rtx-neuro` suite (16 crates): EEG/MEG, LSL real-time, source localization, GNN connectivity |
| Physics-informed | `rtx-neural-operator` (FNO, DeepONet), `rtx-piddm`, `rtx-digital-twin` |
| Scientific ML | `rtx-fea` (FEM), `rtx-cfd` (CFD), `rtx-science` |
| Interpretability | `rtx-interpret` (SAE, attribution, mechanistic interpretability) |
### Meta-Crates (3 user-facing bundles)
- `rtx-core` — Essential functionality
- `rtx-training` — Complete training stack
- `rtx-inference-stack` — Production inference
## Key APIs
```rust
// Tensor API (PyTorch-like)
let x = Tensor::randn([batch, seq_len, hidden], &device)?;
let output = transformer.forward(&x)?;
let loss = cross_entropy(&output, &labels)?;
loss.backward()?;
optimizer.step()?;
// Distributed training
let trainer = DistributedTrainer::new(model, num_gpus)?;
trainer.train(&dataset)?;
// Flash Attention
let attn = FlashAttention::new(head_dim, num_heads)?;
let out = attn.forward(&q, &k, &v, causal=true)?;
// Speculative decoding
let engine = SpeculativeEngine::new(draft_model, target_model)?;
let tokens = engine.generate(&prompt, max_tokens=512)?;
```
## GPU Backend Support
| Backend | Hardware | Feature Flag |
|---------|----------|-------------|
| CUDA | NVIDIA (RTX 4090/5090, A100, H100) | `cuda` |
| Metal | Apple M-series (M1/M2/M3/M4/M5) | `metal` |
| ROCm | AMD RDNA 2/3 (RX 7900 XTX) | `rocm` |
| SYCL | Intel Arc (A770/A750) — **experimental skeleton, ops return NotImplemented** | `sycl` |
| WebGPU | Browser (WASM) | `webgpu` |
| CPU | x86/ARM with MKL | default |
The meta-crates (`rtx-core`, `rtx-training`, `rtx-inference-stack`) expose `cuda`/`metal` features that thread GPU support through their sub-crates.
## Novel Features
- **Speculative decoding** (`rtx-inference::speculative` + `rtx-inference::medusa`/`lookahead`): 23× inference speedup via draft model + target verification — GPU-native Rust implementation
- **Kernel fusion DSL** (rtx-fusion): Eliminates intermediate tensor writes, 50%+ bandwidth savings
- **Neural Operators** (rtx-neural-operator): FNO and DeepONet for mesh-agnostic PDE solving
- **Federated learning at scale** (rtx-federated): 100K+ edge device coordination
- **MRI→FEM pipeline** (rtx-mri2fe): Medical image to finite element mesh for biomechanical simulation
## Platform Dependencies
- **RNCCL** — AllReduce collective operations for distributed training gradient aggregation
- **rustG** — GPU-accelerated compilation toolchain for kernel synthesis
## Feeds Into
- **QPUDIDP** — rustytorch provides the ML training infrastructure for surrogate model training
- **qstar** — the Python qstar surrogate models could be retrained using rustytorch
- Jupyter notebooks via PyO3 Python bindings
## Platform Role
**Layer 8 — ML Infrastructure.** rustytorch is the ML framework that powers the ML-driven components of the QuantumRedClaw platform. QPUDIDP's surrogate models (MLP, normalizing flows, MC Dropout) and the HyperNQ neural decoder in qstar-rs ultimately depend on GPU compute infrastructure. rustytorch provides that infrastructure in pure Rust, matching the platform's no-C/C++ philosophy.
## JEPA Platform (Batches 2026, complete 2026-06-27)
All in `crates/training/rtx-transformers/src/ssl/`. 163 tests passing.
- **I-JEPA** (`jepa.rs`): BlockMaskStrategy, JepaPredictor, JepaTrainer, jepa_loss, EmaTargetEncoder, FeatureBank + k-NN, LinearProbe, JepaEvaluator
- **V-JEPA** (`vjepa.rs`): PatchEmbed3D, TubeMaskStrategy, VJepaTrainer
- **Neuro-JEPA** (`vjepa.rs`): NeuroJepaConfig, NeuroMaskStrategy for EEG/MEG channel-tube masking
- **ViT encoder bridge** (`jepa_vit.rs`): `JepaEncoder` trait, CpuViTEncoder (Tiny→Huge), EmaViTEncoder, JepaTrainerV2 with timing metrics
- **Data pipeline** (`jepa_data.rs`): MultiScaleRandomCrop, JepaAugmentationPipeline, InMemoryShard/WebDatasetShard, JepaBatch, JepaDataPipeline, DatasetStats
- **Cluster config** (`jepa_cluster.rs`): GpuSpec/rtx5060ti, ClusterTopology, JepaParallelConfig (TP/PP/DP auto), CompressionMethod, AdaptiveBatchSizer (GNS-based), ClusterTrainingPlan
**Done (2026-07-09):** GPU backend wiring — full GPU-resident ViT block in `jepa_gpu.rs` (cuBLAS GEMMs + nvrtc layernorm/GELU/softmax/head-slice kernels, CPU↔GPU parity < 1e-3 verified on RTX 5060 Ti); WebDataset `.tar` filesystem reading (`WebDatasetShard::load`); cluster-plan wiring (`JepaRunConfig::apply_cluster_plan`).
**Done (2026-07-10):** multi-node TCP AllReduce (`JepaGradSync::tcp`, wired into the runner for world_size > 1, two-rank loopback training tested); `rtx-jepa` training CLI (`crates/tooling/rtx-jepa-cli`: train/bench/plan/validate); gzip WebDataset shards; full-depth GPU and CPU ViT (all bring-up block caps removed, parity verified); `ViTSizeStr::Micro` test size.
**Done (2026-07-10, round 2):** NCCL GPU-direct AllReduce (`nccl` feature; id bootstrap over the TCP rendezvous; verified single-rank on RTX 5060 Ti); k-NN + linear-probe eval wired into the training loop (results in summary/CSV/CLI); GPU-run checkpointing + eval via `context_encoder_cpu_weights()`.
**Done (2026-07-10, round 3):** GPU weight re-upload on checkpoint resume (restore into host copy + push to device buffers; verified live: 20-step train → resume → 10 further steps).
**Next (needs multi-GPU hardware):** tensor/pipeline parallel execution (config exists, execution is DP-only); multi-rank NCCL verification.
## Current State
Production-ready. 113 crates. Full transformer stack, Flash Attention (v2+v3), MoE, speculative decoding (Medusa/EAGLE/Lookahead), federated learning, NAS, medical/neuroimaging/scientific domain stacks, complete JEPA platform (Batches 2026). 30+ demo applications. 13,000+ tests.
Known honesty notes (2026-07-09 sweep, updated 2026-07-10): `rtx-backend-sycl` is an experimental skeleton (ops return NotImplemented). The formerly-simulated demos now run real compute (`rtx-distllm-demo` real tensors + measured metrics with only the network topology simulated; `rtx-model-zoo` drives a real InferenceEngine with untrained weights and labeled toy output proxies; `rtx-inference-profiler` profiles real matmul/softmax). MoE/flash-attention duplication was audited — see `docs/consolidation.md`: the flagged sites were scaffolding/stubs, not duplicates; `revolutionary/orchestrator_core.rs` execute paths are no-op stubs and `tensor_core_kernels.rs` is a planner with no kernels. Serving/streaming return 503/error until an engine+model is attached; tokenization uses `ServingTokenizer` (HuggingFace tokenizer.json or byte-level fallback).