VITRA trains end-to-end on AMD Instinct MI300X (ROCm) — no fork, no hipify, no custom kernels
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 502
- Forks
- 34
- PR merge metrics
- No merged PRs in 30d
Description
Follow-up to #42, which reported that inference runs on AMD ROCm. This issue covers the other half: end-to-end multi-GPU training.
Summary
I ran VITRA's pretraining loop end-to-end on 8× AMD Instinct MI300X (gfx942) under ROCm and it works with stock upstream code — no fork, no hipify, no custom CUDA/HIP kernels. The training path is pure PyTorch FSDP + RCCL, and neither DeepSpeed nor FlashAttention is needed (the DiT action head uses torch SDPA and the VLM uses the default HF attention).
Getting a clean run required only one runtime env var, one small (platform-agnostic) dataloader fix, and honoring the utils3d pin that's already in pyproject.toml. Details below in case it helps other AMD users, and I'm happy to send a small PR for the code part.
Environment
- GPUs: 8× AMD Instinct MI300X OAM,
gfx942, 192 GB each - Container:
rocm/pytorch:rocm7.2.1_ubuntu24.04_py3.12_pytorch_release_2.9.1(torch2.9.1+rocm7.2.1, HIP 7.2) - Launcher:
torchrun --standalone --nproc_per_node=8 scripts/train.py - Data: a 10k-episode crop of VITRA-1M SSv2 (9,403 webm / 213,835 frames),
data_mix=ssv2 - Config:
human_pretrain.jsonwith short-run overrides for a smoke of the full loop:
batch_size=8(per GPU),total_batch_size=64(→ grad_accum=1),max_steps=1000,
save_steps=250,llm_freeze_step=200, FSDPshard-grad-op(ZeRO-2), bf16.
What worked out of the box
pip install --no-deps -e .+ the pure-Python training deps (transformers 4.47.1, diffusers, timm, peft, etc.). DeepSpeed and FlashAttention were intentionally omitted and are not required.- Model build (PaliGemma2-3B backbone + DiT head), bf16 forward/backward/optimizer step.
- FSDP sharding, checkpoint save and resume (optimizer state included).
apexfused_layer_normJIT-compiled cleanly viahipcc --offload-arch=gfx942("unsupported CUDA calls: 0").
Three changes needed
| # | Symptom | Fix | Type |
|---|---|---|---|
| 1 | First dist.barrier() aborts in RCCL init: [FATAL] HSA_NO_SCRATCH_RECLAIM=1 must be set ... → ncclSystemError |
-e HSA_NO_SCRATCH_RECLAIM=1 |
MI300X env (not code); single-GPU never hits it — worth a docs note |
| 2 | Under many concurrent readers, decord's threaded decoder fails (avcodec_send_packet ... -11), so _grab_window_images never assigns imgs → UnboundLocalError |
decord.VideoReader(name, num_threads=1) in video_utils.py |
Real code fix, platform-agnostic (hits NVIDIA too at high dataloader concurrency) |
| 3 | AttributeError: utils3d.numpy has no attribute 'image_uv' on the augment path |
Honor the utils3d pin already in pyproject.toml (@d790d33) |
Not a VITRA bug — just don't install unpinned |
Only #2 is an actual source change; the files themselves are fine (they decode in isolation at any thread count). Suggest also making _grab_window_images fail-fast after its retries instead of falling through to np.stack(imgs, ...).
Evidence (1000-step run)
What "training works" means here — each of these was exercised on gfx942 with no source changes:
- Full step:
forward → backward → optimizer.step()in bf16, finite loss (no NaN/Inf). - FSDP sharding (
shard-grad-op/ ZeRO-2) across 8 ranks. - RCCL collectives (init + all-gather / reduce-scatter) — the same path that fails without env fix #1.
- Checkpoint save at steps 250 / 500 / 750 / 1000.
- Checkpoint resume — model + optimizer state reloaded correctly and training continued.
- LLM freeze → unfreeze schedule (
llm_freeze_step=200) so the backbone-trainable path is also exercised (backbone LR goes 0 → 1e-5). - 8/8 GPUs at 100% utilization (
rocm-smi), ~10 s/it, clean exit (torchrunexit code 0).
Raw per-step log (rank 0), loss stays finite and trends down (~0.29 → ~0.21):
=>> [Epoch 000] Global Step 000516 =>> Backbone LR :: 0.000000 - Loss :: 0.2664
=>> [Epoch 000] Global Step 000517 =>> Backbone LR :: 0.000000 - Loss :: 0.2943
=>> [Epoch 000] Global Step 000518 =>> Backbone LR :: 0.000000 - Loss :: 0.3123
...
=>> [Epoch 000] Global Step 000813 =>> Backbone LR :: 0.000010 - Loss :: 0.2362
...
=>> [Epoch 000] Global Step 000999 =>> Backbone LR :: 0.000010 - Loss :: 0.2145
(This is a short pipeline-validation run, not a convergence study — the point is the full ROCm training loop runs correctly, not the final metric.)
Offer
Happy to open a small PR with:
decord.VideoReader(name, num_threads=1)+ fail-fast after retries in_grab_window_images, and- a short "Running on AMD ROCm (MI300X)" section in the README (the
HSA_NO_SCRATCH_RECLAIM=1note + the container tag above).
Thanks for open-sourcing VITRA — it's a great codebase to work with.
Contributor guide
No contributing guide indexed for this repository
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Start in video_utils.py at _grab_window_images and inspect the decord.VideoReader call and retry path. Reproduce the threaded-decoder failure with concurrent readers, then verify the decoder fix and fail-fast behavior. If covering the documentation part, update the README with the MI300X ROCm environment note; done means the smoke training run proceeds without the dataloader error.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python
- Domain
- distributed-systems, documentation, machine-learning
- Issue type
- Bug
- Difficulty
- 3/5
- Estimated time
- 1-2 days
- Activity status
- Quiet
- Clarity
- Clearly specified
- Newbie friendliness
- 68/100