Run JAX MaxText training benchmarks with CVS#
2026-09-24
7 min read time
JAX training in CVS is jaxmaxtext. The legacy jax_llama3_1_* suites have
been removed. Use jaxmaxtext_single with a single-node config and
jaxmaxtext_distributed with a distributed config; the mode is inferred from
the config (distributed configs carry the NCCL RDMA device-selection vars in
container.env, single-node configs omit them). Copy the matching
*_threshold.json into the same directory as the suite config.
Set up config#
List available JAX MaxText configuration files:
cvs config list training/jaxmaxtext
Copy the configuration file and the platform threshold file for your GPU, for example:
cvs config copy training/jaxmaxtext/mi3xx_jaxmaxtext_llama-3.3-70b_single.json --output ~/cvs_workspace/training/jaxmaxtext/mi3xx_jaxmaxtext_llama-3.3-70b_single.json cvs config copy training/jaxmaxtext/mi300x_jaxmaxtext_llama-3.3-70b_single_threshold.json --output ~/cvs_workspace/training/jaxmaxtext/mi300x_jaxmaxtext_llama-3.3-70b_single_threshold.json
The
mi3xx_*configs run on both MI300X and MI325X; threshold files stay per-platform (mi300x_*/mi325x_*).Replace every
<changeme>with cluster-specific values: thethreshold_jsonfilename for your GPU platform and thecontainer.imagetag on every config, plus the NCCL/RDMA fields incontainer.envon distributed configs.Change any other parameters relevant to your testing requirements. On MI300X, if a sweep OOMs, lower
per_device_batch_size(theBSin the sweep key) — themi3xx_*configs ship the larger MI325X batch sizes.
Full parameter list: JAX MaxText training configuration files for Cluster Validation Suite (CVS).
Run tests#
JAX MaxText test scripts#
You can list all available JAX MaxText test cases using the CLI:
cvs list jaxmaxtext_single
Available tests in jaxmaxtext_single:
- test_launch_container
- test_setup_tokenizer
- test_smoke
- test_training_run
- test_metric
- test_loss_curve
- test_checkpoint_resume
- test_print_results_table
- test_teardown
cvs list jaxmaxtext_distributed
Available tests in jaxmaxtext_distributed:
- test_launch_container
- test_setup_rdma
- test_setup_tokenizer
- test_smoke
- test_training_run
- test_metric
- test_loss_curve
- test_checkpoint_resume
- test_print_results_table
- test_teardown
Use these scripts to run the JAX MaxText tests.
Single-node:
cvs run jaxmaxtext_single --cluster_file input/cluster_file/cluster.json --config_file input/config_file/training/jaxmaxtext/mi3xx_jaxmaxtext_llama-3.3-70b_single.json --html=/var/www/html/cvs/jaxmaxtext_single.html --capture=tee-sys --self-contained-html --log-file=/tmp/jaxmaxtext_single.log -vvv -s
Distributed:
cvs run jaxmaxtext_distributed --cluster_file input/cluster_file/cluster.json --config_file input/config_file/training/jaxmaxtext/mi3xx_jaxmaxtext_llama-3.3-70b_distributed.json --html=/var/www/html/cvs/jaxmaxtext_distributed.html --capture=tee-sys --self-contained-html --log-file=/tmp/jaxmaxtext_distributed.log -vvv -s
The mode is inferred from the config: distributed configs carry the NCCL RDMA
device-selection vars in container.env (and add the test_setup_rdma
stage), while single-node configs omit them. Run single-node configs with
jaxmaxtext_single and distributed configs with jaxmaxtext_distributed —
the suite must match the config.
Prerequisites#
The following prerequisites are required.
Passwordless SSH from the control host to each node (key in the cluster file), and Docker available on the nodes.
A container image bundling MaxText/JAX for ROCm (config
container.image).A Hugging Face token file at
paths.hf_token_file(used to fetch the tokenizer; needs network access on the nodes). Skipped when every enabled run usesdataset_type: synthetic.A shared filesystem (
paths.shared_fs) reachable from all nodes for the models cache and logs.
Test lifecycle#
The tests run in this fixed order. [sweep] = one row per enabled sweep;
[sweep-metric] = one row per metric per sweep.
Order |
Test |
Runs on |
Purpose |
|---|---|---|---|
1 |
|
once |
Launch and verify the container (and check out |
2 |
|
distributed only |
Copy the RDMA lib into the container (thor2 NIC) and verify |
3 |
|
once |
Download the HF tokenizer. Skipped when every enabled run is |
4 |
|
once |
Small fixed run (BF16, few steps): the model loads and trains with no error/NaN signature. Enabled by default; a failure gates the rest of the suite. Skip with |
5 |
|
per sweep |
Build the command, train, poll, and parse results. |
6 |
|
per sweep × metric |
Threshold PASS/FAIL per metric. |
7 |
|
per sweep |
Render the loss PNG and gate on a downward trend. |
8 |
|
once |
Opt-in ( |
9 |
|
once |
Console tables + metric-results HTML + failure summary. |
10 |
|
once |
Tear the container down. |
A training failure is isolated to that sweep’s test_training_run row; other
sweeps still run, and the failed sweep’s downstream test_metric /
test_loss_curve rows are skipped. Lingering ranks are killed before the next
sweep launches.
Sweeps#
A sweep is one full training run. sweeps is a {key: overrides} map
and runs selects which sweep keys to execute. The key is a parseable spec —
BS=<batch>,PRECISION=<BF16|FP8>,SL=<seqlen> (e.g. BS=3,PRECISION=FP8,SL=8192)
— that CVS turns into maxtext_config overrides (per_device_batch_size /
quantization / max_target_length); add any extra override (e.g. a per-sweep
steps) inside the sweep’s {}. The key is also the threshold cell key and
appears in every parametrized row.
Metrics and PASS/FAIL#
Each test_metric[sweep-metric] compares the parsed metric against its spec
in the sweep’s cell of the threshold file:
Status |
Meaning |
|---|---|
PASS |
Value satisfies the threshold. |
FAIL |
Value violates the threshold (row is red; aggregated in the summary). |
N/A |
Metric not produced this run (feature disabled / rampup) — not a failure. |
RECORD |
No threshold, or |
Metrics use the training. namespace (tflops_per_sec_per_gpu,
tokens_per_sec_per_gpu, tokens_per_sec_total, scaling_efficiency_pct,
step_time_*, final_loss, loss_decreased, eval_loss,
steps_to_target, time_to_target_seconds). Gating requires
enforce_thresholds: true; an info threshold always passes (record-only).
See JAX MaxText training configuration files for Cluster Validation Suite (CVS) for threshold kinds
and the full metric list.
Reports and logs#
Results table — one row per test; metric rows show PASS/FAIL.
Full Log — each test row links to its own captured log.
Metric Results — every
test_metricrow links to a sharedmetric_results.html(Sweep | Metric | Expected | Actual | Unit | Status).Loss Curve — each
test_loss_curverow links to a per-sweep PNG.Console summary —
test_print_results_tableprints per-sweep tables and an aggregated list of failed(sweep, metric)checks.
Training-log error detection#
During polling, each node’s training.log is scanned for the regexes in
error_patterns (config-driven; falls back to built-in defaults
covering NCCL, GPU HW faults, assertion/JAX stack traces, ROCm init errors,
Python fatal errors, TF coordination errors, RESOURCE_EXHAUSTED/OOM, and
segfaults). Fatal Python crashes (tracebacks / import errors) are always scanned
so an early failure fails fast instead of running to the poll timeout. A match
fails that sweep’s test_training_run with the matched signature name.