JAX MaxText training configuration files for Cluster Validation Suite (CVS)#
2026-09-24
18 min read time
The JAX MaxText suites (jaxmaxtext_single / jaxmaxtext_distributed) run
MaxText pre-training inside a
container on one or more nodes and gate the run on performance and correctness
metrics with a PASS/FAIL HTML report.
Note
JAX training in CVS is jaxmaxtext. The legacy jax suites
(jax_llama3_1_*) have been removed; use jaxmaxtext_single /
jaxmaxtext_distributed.
The JAX MaxText tests check:
Container orchestration: Docker setup with ROCm/RDMA.
Model load + smoke: The model loads and trains a few steps with no error/NaN signature.
Per-sweep training: One full run per enabled sweep (e.g., BF16, FP8).
Performance targets: TFLOP/s, tokens/s, step time, and multi-node scaling efficiency.
Convergence: Final loss / loss-decreasing trend, optional time-to-target.
Checkpoint save/resume (opt-in): resume correctness + checkpoint I/O timing.
Use cvs config list training/jaxmaxtext to list available templates, or
cvs config copy training/jaxmaxtext/<name> to copy one to your working directory.
Note
Any value containing
<changeme>must be replaced for your setup.container.imageships with a<changeme>tag (e.g.rocm/jax-training:maxtext-v26.4 <changeme>) — set the image tag your cluster has before running; unlike the env placeholders it is not caught at config load but fails at container launch until replaced. Distributed configs additionally ship the NCCL RDMA/NIC device-selection vars incontainer.envwith an example value plus a<changeme>tag (NCCL_IB_HCA,NCCL_SOCKET_IFNAME,GLOO_SOCKET_IFNAME,NCCL_IB_GID_INDEX); the config hard-exits at load while anycontainer.envvalue still contains<changeme>.{user-id}resolves to the cluster/OS username at runtime, and{shared_fs}/{paths.*}self-references resolve from thepathsblock.Keys prefixed with
_(e.g._env_comment,_train_params_comment) are inline comments and are ignored by the loader.
The suite/lifecycle reference is in
Run JAX MaxText training benchmarks with CVS. The config files themselves live in
cvs/input/config_file/training/jaxmaxtext/ (each config plus a sibling
_threshold.json); this page documents every block and the threshold format.
Available configurations#
Config files follow the naming pattern <gpu>_jaxmaxtext_<model>_<mode>.json.
Each config has a sibling _threshold.json referenced by threshold_json.
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); single-node configs omit them. Run single-node configs with
jaxmaxtext_single and distributed configs with jaxmaxtext_distributed.
The mi3xx_* configs are GPU-generic and run on both MI300X and MI325X
(both CDNA3). They ship the MI325X batch sizes, so on MI300X (less HBM) a large
batch may hit OOM — lower per_device_batch_size (the BS in the sweep
key) if you see one. Because threshold targets are platform-specific, each
mi3xx_* config ships threshold_json with a <changeme> tag: point it at
the mi300x_* or mi325x_* _threshold.json for your GPU before
running (the config otherwise refuses to load). The mi35x_* config targets
MI350-class (CDNA4) GPUs.
Config file |
GPU(s) |
Mode |
Precisions |
|---|---|---|---|
|
MI300X / MI325X |
distributed |
BF16, FP8 |
|
MI300X / MI325X |
single |
BF16, FP8 |
|
MI300X / MI325X |
distributed |
BF16, FP8 |
|
MI300X / MI325X |
single |
BF16, FP8 |
|
MI300X / MI325X |
distributed |
BF16, FP8 |
|
MI300X / MI325X |
distributed |
BF16 |
|
MI300X / MI325X |
distributed |
BF16, FP8 |
|
MI300X / MI325X |
distributed |
BF16 |
|
MI350-class |
single |
BF16, FP8 |
Threshold files stay per-platform (mi300x_*_threshold.json and
mi325x_*_threshold.json); the mi3xx_* config selects one via
threshold_json.
Config layout#
A config groups its keys into five areas:
CVS params at the root —
gpu_name,gpus_per_node,threshold_json,enforce_thresholds, andpaths.container — image, Docker runtime args, and the static
envexported into the container on every node.train_params — the tokenizer source, train-script candidates, the
maxtext_configpassthrough written verbatim to the MaxText YAML, and the structuredxla_flags(exported as oneXLA_FLAGSenv var).tests blocks at the root —
scaling_baseline,convergence,loss_curve,smoke,checkpoint_resume,error_patterns.sweeps + runs —
sweepsis a{key: overrides}map (one full training run each) whose key encodes the primary params (BS=..,PRECISION=..,SL=.., parsed by CVS);runsselects which sweep keys to execute.
Example configuration#
A representative distributed config
(mi3xx_jaxmaxtext_llama-3.3-70b_distributed.json, abridged):
mi3xx_jaxmaxtext_llama-3.3-70b_distributed.json (abridged)
{
"gpu_name": "mi3xx",
"threshold_json": "mi325x_jaxmaxtext_llama-3.3-70b_distributed_threshold.json <changeme>",
"enforce_thresholds": false,
"gpus_per_node": 8,
"paths": {
"shared_fs": "/home/{user-id}",
"models_dir": "{shared_fs}/cache/maxtext",
"log_dir": "{shared_fs}/LOGS/jaxmaxtext",
"hf_token_file": "{shared_fs}/.hf_token",
"temp_dir": "/tmp/{user-id}/jaxmaxtext"
},
"container": {
"lifetime": "per_run",
"name": "rocm-jaxmaxtext-llama3.3-70b",
"image": "rocm/jax-training:maxtext-v26.4 <changeme>",
"runtime": { "name": "docker", "args": { "network": "host", "ipc": "host", "privileged": true, "shm-size": "256G", "ulimit": ["nofile=65535:65535"], "volumes": ["..."] } },
"env": {
"GPU_MAX_HW_QUEUES": "2",
"HSA_FORCE_FINE_GRAIN_PCIE": "1",
"XLA_PYTHON_CLIENT_MEM_FRACTION": "0.97",
"NCCL_IB_HCA": "rdma0,rdma1,rdma2,rdma3,rdma4,rdma5,rdma6,rdma7 <changeme>",
"NCCL_SOCKET_IFNAME": "eno0 <changeme>",
"GLOO_SOCKET_IFNAME": "eno0 <changeme>",
"NCCL_IB_GID_INDEX": "3 <changeme>",
"NCCL_IB_DISABLE": "0",
"NCCL_IB_TC": "41",
"NCCL_IB_SL": "0",
"NVTE_FUSED_ATTN": "1",
"JAX_COORDINATOR_PORT": "12346",
"JAX_DISTRIBUTED_INITIALIZATION_TIMEOUT_SECONDS": "1800",
"JAX_DISTRIBUTED_HEARTBEAT_TIMEOUT_SECONDS": "900"
}
},
"train_params": {
"hf_model_id": "NousResearch/Meta-Llama-3-70B",
"train_script_paths": ["/workspace/maxtext/src/maxtext/trainers/pre_train/train.py", "/workspace/maxtext/src/MaxText/train.py"],
"maxtext_config": {
"base_config": "base.yml",
"model_name": "llama3.3-70b",
"tokenizer_path": "{paths.models_dir}/Meta-Llama-70-B",
"hardware": "gpu",
"steps": 30,
"enable_checkpointing": false,
"attention": "cudnn_flash_te",
"dtype": "bfloat16",
"weight_dtype": "bfloat16",
"dataset_type": "synthetic",
"quantization": "",
"per_device_batch_size": 3,
"max_target_length": 8192,
"remat_policy": "full",
"scan_layers": true,
"ici_fsdp_parallelism": 8,
"dcn_data_parallelism": -1
},
"xla_flags": { "xla_gpu_autotune_level": "0", "...": "..." }
},
"scaling_baseline": { "tokens_per_sec_total": 394000.0, "num_nodes": 1 },
"convergence": { "target_metric": "auto", "target_value": 10.0 },
"loss_curve": { "sample_every": 10, "milestone_steps": [100, 500, 1000, 5000], "max_slope": 0.0, "enforce": true },
"smoke": { "enabled": true, "steps": 5, "per_device_batch_size": 1, "max_target_length": 2048 },
"checkpoint_resume": { "enabled": false, "steps_before_ckpt": 6, "steps_after_resume": 6, "checkpoint_period": 5, "loss_tolerance": 0.1, "delete_ckpt_dir": true },
"error_patterns": { "NCCL ERROR": "NCCL ERROR|NCCL timeout", "...": "..." },
"sweeps": {
"BS=3,PRECISION=BF16,SL=8192": { "_comment": "extra maxtext_config overrides go here, e.g. \"steps\": 300" },
"BS=3,PRECISION=FP8,SL=8192": { "_comment": "extra maxtext_config overrides go here, e.g. \"steps\": 300" }
},
"runs": ["BS=3,PRECISION=BF16,SL=8192", "BS=3,PRECISION=FP8,SL=8192"]
}
Top-level (CVS) fields#
The following fields sit at the root of the configuration object and control suite-level behavior.
Field |
Example |
Description |
|---|---|---|
|
|
GPU architecture label (informational; also used in run/report labels).
|
|
|
GPUs per node; |
|
|
If |
|
|
Companion threshold filename, resolved next to the config. The
|
paths#
Field |
Example |
Description |
|---|---|---|
|
|
Base path reachable from all nodes. Self-referenced by other |
|
|
Tokenizer/model cache directory. |
|
|
Training log output directory (per-node logs are namespaced under it). |
|
|
Hugging Face token file (for the tokenizer download). |
|
|
Host-user-namespaced in-container scratch for launcher scripts / MaxText YAML. Keep |
container#
Field |
Example |
Description |
|---|---|---|
|
|
Launched once per session, torn down after. |
|
|
Container instance name (any unique string). |
|
|
Required — the MaxText/JAX ROCm image present on all nodes. Ships with a |
|
(see snippet) |
Docker args: |
|
(dict) |
Static environment exported into the container on every node at |
container.env#
A flat {NAME: value} map exported into the container on every node at
docker run time and inherited by docker exec (there is no env script to
source). Per-node/dynamic vars (JAX_COORDINATOR_IP, NNODES,
NODE_RANK, JAX_PROCESS_INDEX) and credentials (HF_TOKEN, HF_HOME,
LD_LIBRARY_PATH, PYTHONPATH) are injected by the launcher and must not
be set here. XLA_FLAGS is not set here either — it is built from the
structured train_params.xla_flags map. The vars are grouped for readability:
Group |
Examples |
Description |
|---|---|---|
GPU / memory |
|
ROCm/HIP tuning and the fraction of GPU memory JAX may allocate (e.g. |
NCCL device selection (distributed) |
|
Cluster-specific RDMA/NIC device selection, shipped as |
NCCL tuning |
|
RoCE/IB transport tuning ( |
Transformer-Engine / Composable-Kernel |
|
Fused-attention numerics controls (numerics-sensitive; tune per model/BKC). |
JAX coordinator |
|
JAX distributed coordinator port and the init-rendezvous / heartbeat timeouts. |
Note
NCCL device selection (distributed). Each of
NCCL_IB_HCA, NCCL_SOCKET_IFNAME, GLOO_SOCKET_IFNAME, and
NCCL_IB_GID_INDEX ships with an example value followed by a <changeme>
tag (e.g. "rdma0,...,rdma7 <changeme>", "eno0 <changeme>",
"3 <changeme>"). Replace the whole value with your cluster’s setting — the
config hard-exits at load while any container.env value still contains
<changeme>. Discover them with ibv_devices (HCAs) and ip -br link
(host interface). GLOO_SOCKET_IFNAME accepts a single interface only.
train_params#
Field |
Example |
Description |
|---|---|---|
|
|
Hugging Face repo the tokenizer is downloaded from. Download is skipped when every enabled run uses |
|
(list) |
Candidate in-container MaxText entrypoints; the job picks the first that exists (list newest-first, e.g. the v26.4+ path before the v26.3 path). |
|
|
Optional CVS-side label for run names / report / loss-curve filenames. Defaults to |
|
|
Optional. When set, after container launch the job runs |
|
|
MaxText repository root inside the container (used by the branch checkout / install command). |
|
|
Optional shell command run verbatim from |
train_params.maxtext_config#
Written verbatim into the MaxText YAML, so any valid MaxText parameter can be
set here (including steps, enable_checkpointing, and tokenizer_path).
run_name and base_output_directory are injected by the driver and must
not be set here. The most-edited keys:
Key |
Example |
Description |
|---|---|---|
|
|
MaxText base config to inherit defaults from. |
|
|
MaxText model preset (layers/heads/dims). Must exist in the image’s MaxText. |
|
|
In-container directory the tokenizer is written to / read from (the download target). |
|
|
Target backend. |
|
|
Training steps; also drives completion detection and the poll budget. |
|
|
Whether MaxText writes checkpoints during the run. |
|
|
Attention kernel: |
|
|
Compute dtype and master-weight dtype. |
|
|
|
|
|
|
|
|
Per-GPU batch and sequence length (typically overridden per sweep). |
|
|
Activation rematerialization (memory vs. recompute trade-off). |
|
|
Scan the decoder stack (memory/compile savings). |
|
|
Intra-node parallelism dims (fsdp/data/tensor/sequence/pipeline/expert). |
|
|
Cross-node parallelism dims; |
Other passthrough keys seen in the configs: packing, megablox /
sparse_matmul / capacity_factor / sharding_tolerance (MoE kernel
path), profiler / skip_first_n_steps_for_profiler / profiler_steps,
shardy, logits_dot_in_fp32, param_scan_axis, max_segments_per_seq,
kv_quant_*, optimizer_memory_host_offload, async_checkpointing,
log_period, enable_goodput_recording / monitor_goodput.
Note
Using real data (HuggingFace). Configs default to dataset_type:
synthetic (random tokens; no data/tokenizer download) for throughput and
functional runs. For a genuine loss curve, set these keys inside
maxtext_config (the HF pipeline streams data — no full download):
"dataset_type": "hf",
"hf_path": "allenai/c4",
"hf_data_dir": "en",
"train_split": "train",
"tokenizer_type": "huggingface"
tokenizer_type: huggingface is required so MaxText loads the HF
tokenizer.json named by hf_model_id (the sentencepiece default
would mismatch it), and real data triggers the tokenizer download. Do not
place comment (_-prefixed) keys inside maxtext_config — every key there
is written verbatim to the run YAML and MaxText rejects unknown keys.
train_params.xla_flags#
A structured {flag: value} map emitted as a single XLA_FLAGS env var
(--<flag>=<value> ...) into container.env. An empty map omits
XLA_FLAGS entirely (so XLA’s own defaults are not clobbered). Notable entries
include xla_gpu_autotune_level, xla_gpu_enable_latency_hiding_scheduler,
and the all-gather / reduce-scatter combine thresholds.
Tests blocks#
This section describes the optional test blocks that extend the base training run with additional validation passes.
scaling_baseline (distributed)#
Field |
Default |
Description |
|---|---|---|
|
|
Single-node total tok/s baseline ( |
|
|
Nodes used to produce the baseline ( |
convergence#
Field |
Default |
Description |
|---|---|---|
|
|
|
|
|
Loss target for |
loss_curve#
Field |
Default |
Description |
|---|---|---|
|
|
Sample training loss every N steps for the slope check. |
|
|
Steps always included in the sampled curve. |
|
|
Least-squares slope must be |
|
|
|
smoke#
The smoke test (test_smoke) loads the model and runs a few steps at a small
fixed batch/seqlen in BF16, passing only if no error/NaN signature fires (no
metric checks). A failure gates the rest of the suite. When the block is omitted,
the schema defaults apply (enabled).
Field |
Default |
Description |
|---|---|---|
|
|
Runs by default (opt-OUT). Set |
|
|
Steps for the smoke run. |
|
|
Small fixed batch. |
|
|
Small fixed sequence length. |
checkpoint_resume#
Opt-in (enabled: false). Runs one sweep twice: Phase 1 trains
steps_before_ckpt with checkpointing on (saved at checkpoint_period);
Phase 2 resumes and trains steps_after_resume more. Passes when Phase 2
restarts at the checkpoint step and the boundary loss matches Phase 1 within
loss_tolerance.
Field |
Default |
Description |
|---|---|---|
|
|
Opt-in switch. |
|
|
Which sweep to exercise ( |
|
|
Phase-1 steps (checkpoint saved at |
|
|
Phase-2 steps after resuming. |
|
|
Save frequency; must be |
|
|
Max loss delta at the resume boundary. |
|
|
I/O time gates for |
|
|
Delete the checkpoint dir after the test ( |
|
|
Optional shrink of the model (same tokenizer/vocab) for a fast I/O check. |
error_patterns#
A {name: regex} dict scanned in each node’s training.log during polling;
a match fails that sweep’s test_training_run with the matched name. Remove
the block to use the built-in defaults (NCCL, GPU HW faults, assertion/JAX stack
traces, ROCm init errors, Python fatal errors, TF coordination errors,
RESOURCE_EXHAUSTED/OOM, and segfaults).
Sweeps and runs#
sweeps is a {key: overrides} map — each entry is one full training run,
and its key is both the parsed sweep spec and the threshold cell key.
runs is the list of sweep keys to actually execute.
The key is a comma-separated, parseable spec that CVS turns into
maxtext_config overrides for that run:
Token |
Maps to |
Notes |
|---|---|---|
|
|
Per-GPU batch size. |
|
|
|
|
|
Sequence length. |
dtype / weight_dtype are always bfloat16 (the maxtext_config
default) and are no longer repeated per sweep. Add any extra override inside
the sweep’s {} (it takes precedence over the parsed key), e.g. a per-sweep
steps — which CVS also uses for the run’s timeout/poll budget and completion
detection.
"sweeps": {
"BS=3,PRECISION=FP8,SL=8192": {
"_comment": "extra maxtext_config overrides go here, e.g. \"steps\": 300"
}
},
"runs": ["BS=3,PRECISION=FP8,SL=8192"]
Threshold files#
Each config has a sibling <config-stem>_threshold.json referenced by
threshold_json. It maps each sweep key (cell key) to a dict of
{metric: spec}, one spec per line. A metric is gated (PASS/FAIL) only when
enforce_thresholds: true and it has a numeric spec whose kind is not
info; otherwise it is recorded. The cell key must match the sweep key
exactly, or the metric falls back to RECORD. Metrics not produced by a run
report N/A (not a failure).
Note
The shipped threshold values were captured on a 2-node (2N) run
(single-node configs: 1N; the large llama-3.1-405b and
deepseek-v4-284b configs: 4N). Throughput and step-time scale with the
GPU count, so if you run a different number of nodes, update the values to the
appropriate targets for that node count (see the _node_count_comment in
each threshold file).
"BS=3,PRECISION=BF16,SL=8192": {
"training.tflops_per_sec_per_gpu": {"kind": "min", "value": 260.0},
"training.tokens_per_sec_per_gpu": {"kind": "min", "value": 1217.0},
"training.final_loss": {"kind": "max", "value": 15.0},
"training.loss_decreased": {"kind": "min", "value": 1},
"training.step_time_p95_ms": {"kind": "info", "value": 3600000.0}
}
Threshold kinds#
Each threshold entry specifies how the measured value is compared against the expected value.
Kind |
Passes when |
Notes |
|---|---|---|
|
|
Lower bound. |
|
|
Upper bound. |
|
|
Upper bound, |
|
|
Lower bound, |
|
|
Needs a |
|
|
Needs a |
|
always |
Record-only; keeps a default |
Tracked metrics#
All metrics use the training. namespace.
The suite records the following metrics; each can be referenced in threshold entries by its full dotted name.
Metric |
Single |
Dist. |
Description |
|---|---|---|---|
|
✓ |
✓ |
TFLOP/s per GPU (typically |
|
✓ |
✓ |
Tokens/s per GPU (typically |
|
✓ |
✓ |
Total tokens/s across all GPUs. |
|
— |
✓ |
Multi-node scaling efficiency % vs. |
|
✓ |
✓ |
Mean step wall time (s). |
|
✓ |
✓ |
Step-time mean / p50 / p95 (ms). |
|
✓ |
✓ |
Final training loss (typically |
|
✓ |
✓ |
|
|
✓ |
✓ |
Final eval loss (only when eval is enabled). |
|
✓ |
✓ |
Convergence metrics (only when |
To start gating a metric currently marked info: replace "kind": "info"
with min / max / etc. and set a calibrated value. Checkpoint I/O
timings (checkpoint_save_seconds / checkpoint_load_seconds) are gated by
the checkpoint_resume block, not the threshold file.