jburtoft's picture
docs: add revision=v1.0.0 + trust_remote_code=True to get_kernel() calls
2a17ff8
|
Raw
History Blame Contribute Delete
3.93 kB
# Examples
Runnable samples demonstrating the `mamba3-neuron-kernels` package. Each script
is self-contained and prints its own results.
## Requirements
- **AWS Neuron trn2.3xlarge** instance (or larger)
- **Neuron SDK Beta 3+**: torch-neuronx 2.11.3+, NKI 0.4.0+, PyTorch 2.11+
- **`kernels` library**: `pip install kernels>=0.15.2`
- **Environment**: `export TORCH_NEURONX_ENABLE_CONCATENATION=1` (recommended, ~3-8% speedup)
## Running
Each script prints its own status. If run from a local clone of this repo, the
scripts import from `../build/torch-neuron/`. When the repo is loaded via
`kernels.get_kernel("jburtoft/mamba3-neuron-kernels", revision="v1.0.0", trust_remote_code=True)`, the same code works.
```bash
# One-time setup (on your trn2 instance)
source $HOME/workspace/native_venv/bin/activate
export TORCH_NEURONX_ENABLE_CONCATENATION=1
# Run any sample
cd examples/
python 01_smoke_test.py
python 02_forward_parity.py
python 03_decode_generation.py
python 04_backward_training.py
python 05_benchmark.py # short version (seqlen 64, 256)
python 05_benchmark.py --full # includes seqlen 512, 1024
```
## What each sample does
| Script | Purpose | Runtime |
|---|---|---|
| `01_smoke_test.py` | Loads the kernel, instantiates a mixer, runs one forward. First-run compile takes ~15 s; warm call is a few ms. | ~30 s |
| `02_forward_parity.py` | Runs SISO and MIMO forward with and without the NKI kernel on identical weights, compares outputs. Both should agree at cos_sim > 0.999. | ~60 s |
| `03_decode_generation.py` | Prefills a 64-token context, then runs 128 single-token decode steps. Compares NKI vs eager decode at each step and reports cache correctness + per-step latency. | ~120 s |
| `04_backward_training.py` | Trains SISO and MIMO mixers for 10 steps each on random data with MSE loss. Verifies loss decreases (proves autograd works end-to-end). | ~90 s |
| `05_benchmark.py` | Reproduces the perf tables in the README. Reports forward latency for both modes across seqlens, plus decode step time. | ~90 s short / ~5 min full |
## Expected output (as of v1.0.0)
Target dims: `d_model=1024, d_state=128, headdim=64, nheads=32`, batch=1, trn2.3xlarge LNC=2.
### Forward (with `TORCH_NEURONX_ENABLE_CONCATENATION=1`)
| Path | seqlen=64 | seqlen=256 |
|---|---|---|
| SISO + torch.compile + NKI SSD | ~1.7 ms | ~3.6 ms |
| MIMO + NKI SSD (eager) | ~6.4 ms | ~10.3 ms |
### Decode step (single-token)
| Path | median | p95 |
|---|---|---|
| SISO + NKI SSD | ~4.0 ms | ~4.5 ms |
| MIMO + NKI SSD | ~4.5 ms | ~5.5 ms |
### Batched throughput (batch=32, seqlen=256, eager + NKI SSD)
| Path | per-sample latency |
|---|---|
| SISO | ~1.78 ms/sample |
| MIMO | ~5.76 ms/sample |
## Troubleshooting
**`RuntimeError: NKI kernel requires chunk_size=64`** or similar assertion in `NeuronMamba3Mixer.__init__`:
The NKI kernels are compiled against the reference target config. Use these exact values:
- SISO: `chunk_size=64, mimo_rank=1`
- MIMO: `chunk_size=16, mimo_rank=4`
- Both require: `d_state=128, headdim=64`
For non-reference configs, use `use_nki_ssd=False` to fall back to the eager path.
**First-call latency is very slow (>15 seconds)**:
That's the initial NKI compile. Every unique input shape triggers a fresh compile. Subsequent calls with the same shape use the cached NEFF.
**MIMO forward with `torch.compile` is much slower than eager**:
Known regression -- see [ticket 61](https://github.com/aws-neuron/aws-neuron-sdk/) filing. Don't wrap the MIMO mixer with `torch.compile` until this is fixed. Use eager only for MIMO.
**Silent gradient bias in MIMO training**:
The mixer's forward internally uses `torch.cat` instead of `F.pad` for the beta-term shift, to work around a Neuron autograd bug (`F.pad + einsum` chain produces 2% gradient error on device). This is a workaround the mixer applies transparently. See the `kernels/mixer.py` comment for details.