[AMD] Enable DeepSeek-V4-Flash FP8 RL training on MI355X (#1607)

Co-authored-by: Zhiyao Jiang <jessicajiang324@gmail.com>
This commit is contained in:
Xinyu Jiang
2026-07-09 19:49:03 -07:00
committed by GitHub
co-authored by Zhiyao Jiang
parent 23335971f9
commit 352b4aafcb
5 changed files with 653 additions and 2 deletions
+11 -1
View File
@@ -58,12 +58,14 @@ ENV CMAKE_PREFIX_PATH=/opt/rocm:/opt/rocm/hip:/usr/local:/usr
COPY docker/amd_patch/latest/megatron.patch /tmp/amd_patch/megatron.patch
COPY docker/amd_patch/latest/miles.patch /tmp/amd_patch/miles.patch
COPY docker/amd_patch/latest/sglang_attn_bridge /tmp/amd_patch/sglang_attn_bridge
COPY docker/amd_patch/latest/tile_kernels.patch /tmp/amd_patch/tile_kernels.patch
COPY docker/amd_patch/latest/tilelang_hip_fp8.patch /tmp/amd_patch/tilelang_hip_fp8.patch
COPY requirements.txt /tmp/requirements.txt
# ======================================== Apt dependencies =====================================
RUN apt update
# Install build tools and diagnostics utilities.
RUN apt install -y build-essential cmake dnsutils ethtool git nvtop rsync
RUN apt install -y build-essential cmake dnsutils ethtool git nvtop patch rsync
# Build rccl-tests diagnostics binaries.
RUN git clone --depth 1 --branch ${RCCL_TESTS_BRANCH} ${RCCL_TESTS_REPO} /tmp/rocm-systems && \
@@ -90,6 +92,14 @@ RUN GPU_ARCHS=${GPU_ARCH} BUILD_TARGET=rocm MAX_JOBS=${MAX_JOBS} \
RUN pip install flash-linear-attention==0.4.2
# Enable deepseek-ai/TileKernels on ROCm/gfx950 for mHC and FP8 QAT cast-back.
# These temporary build-time patches make tile_kernels import lazy, remove CUDA/Hopper-only PDL from mhc_post,
# and fix TileLang's HIP fp8->float conversion to use ROCm device intrinsics instead of the SDK host-only operator.
# Drop these patches once the corresponding tilelang/tile_kernels fixes land in the base image.
RUN pip install --no-deps tile_kernels==1.0.0 && \
patch -p1 -d /opt/venv/lib/python3.10/site-packages < /tmp/amd_patch/tile_kernels.patch && \
patch -p1 -d /opt/tilelang/src/tl_templates/hip < /tmp/amd_patch/tilelang_hip_fp8.patch
RUN rm -rf /root/Megatron-LM && \
git clone --recursive -b ${MEGATRON_BRANCH} https://github.com/${MEGATRON_REPO}.git /root/Megatron-LM && \
cd /root/Megatron-LM && \
@@ -0,0 +1,36 @@
diff -ru a/tile_kernels/__init__.py b/tile_kernels/__init__.py
--- a/tile_kernels/__init__.py
+++ b/tile_kernels/__init__.py
@@ -1,15 +1,3 @@
-import tilelang
-from . import (
- config,
- engram,
- mhc,
- modeling,
- moe,
- quant,
- transpose,
- torch,
- testing,
-)
from .config import get_num_sms, get_device_num_sms, set_num_sms
diff -ru a/tile_kernels/mhc/post_kernel.py b/tile_kernels/mhc/post_kernel.py
--- a/tile_kernels/mhc/post_kernel.py
+++ b/tile_kernels/mhc/post_kernel.py
@@ -39,7 +39,6 @@
c_local = T.alloc_fragment(mhc, T.float32)
T.copy(a[pid_n, 0, 0], a_local)
T.copy(c[pid_n, 0], c_local)
- T.pdl_sync()
for i0_h in T.Pipelined(T.ceildiv(h, h_blk), num_stages=2):
T.copy(b[pid_n, 0, i0_h * h_blk], b_shared, disable_tma=True)
diff -ru a/tile_kernels/modeling/__init__.py b/tile_kernels/modeling/__init__.py
--- a/tile_kernels/modeling/__init__.py
+++ b/tile_kernels/modeling/__init__.py
@@ -1,2 +0,0 @@
-from . import engram
-from . import mhc
@@ -0,0 +1,32 @@
--- a/hip_fp8.h
+++ b/hip_fp8.h
@@ -82,7 +82,13 @@
return *reinterpret_cast<const hip_fp8_e4_t *>(&data);
}
__device__ operator float() const {
- return static_cast<float>(static_cast<hip_fp8_e4_t>(*this));
+ constexpr __hip_fp8_interpretation_t interp =
+#if (TILELANG_FP8_E4M3_VARIANT == TILELANG_FP8_E4M3_VARIANT_FNUZ)
+ __HIP_E4M3_FNUZ;
+#else
+ __HIP_E4M3;
+#endif
+ return __half2float(__half(__hip_cvt_fp8_to_halfraw(data, interp)));
}
};
@@ -111,7 +117,13 @@
return *reinterpret_cast<const hip_fp8_e5_t *>(&data);
}
__device__ operator float() const {
- return static_cast<float>(static_cast<hip_fp8_e5_t>(*this));
+ constexpr __hip_fp8_interpretation_t interp =
+#if (TILELANG_FP8_E5M2_VARIANT == TILELANG_FP8_E5M2_VARIANT_FNUZ)
+ __HIP_E5M2_FNUZ;
+#else
+ __HIP_E5M2;
+#endif
+ return __half2float(__half(__hip_cvt_fp8_to_halfraw(data, interp)));
}
};
// Note: E8M0 types are not supported in current HIP version
@@ -17,7 +17,11 @@ class _BFloat16LinearFP32Func(torch.autograd.Function):
ctx.weight_dtype = weight.dtype
x_2d = x_bf16.reshape(-1, x_bf16.shape[-1])
out = torch.mm(x_2d, weight_bf16.t(), out_dtype=torch.float32)
if torch.version.hip is not None:
# ROCm lacks bf16-in/fp32-out torch.mm, so upcast bf16-rounded inputs to fp32 to match CUDA/cublas fp32 accumulation.
out = torch.mm(x_2d.float(), weight_bf16.t().float())
else:
out = torch.mm(x_2d, weight_bf16.t(), out_dtype=torch.float32)
return out.view(*x.shape[:-1], weight_bf16.shape[0])
@staticmethod
+569
View File
@@ -0,0 +1,569 @@
"""
DeepSeek V4 training script.
Supports:
- DeepSeek-V4-Flash-FP8 Public FP8 repackage of deepseek-ai/DeepSeek-V4-Flash
(sgl-project/DeepSeek-V4-Flash-FP8, 291B, 43 layers).
Verified full-model profile: 4 nodes x 8 GPUs on MI355X (gfx950).
- DeepSeek-V4-Flash-FP8-4layer 4-layer prune of the above for single-node
smoke testing. **Cannot generate meaningful output -
pipeline-only sanity check.**
Usage patterns:
1. One-shot full pipeline (download + convert + train):
python scripts/run_deepseek_v4.py full-train \
--model-name DeepSeek-V4-Flash-FP8-4layer \
--num-nodes 1 --num-gpus-per-node 8
2. Individual steps (download -> FP8->BF16 -> BF16->torch_dist -> rsync -> train):
python scripts/run_deepseek_v4.py prepare-download --model-name DeepSeek-V4-Flash-FP8
python scripts/run_deepseek_v4.py prepare-single --model-name DeepSeek-V4-Flash-FP8 \
--hf-checkpoint /root/models/DeepSeek-V4-Flash-FP8
python scripts/run_deepseek_v4.py prepare-spmd --model-name DeepSeek-V4-Flash-FP8 \
--num-nodes 1 --num-gpus-per-node 8
python scripts/run_deepseek_v4.py prepare-cp --model-name DeepSeek-V4-Flash-FP8
python scripts/run_deepseek_v4.py train --model-name DeepSeek-V4-Flash-FP8 \
--num-nodes 4 --num-gpus-per-node 8 \
--hf-checkpoint /root/models/DeepSeek-V4-Flash-FP8
"""
from dataclasses import dataclass, field
from pathlib import Path
from typing import Literal
import typer
import miles.utils.external_utils.command_utils as U
app = typer.Typer()
_DEFAULT_MODEL_ORG = {
"DeepSeek-V4-Flash-FP8": "sgl-project",
# 4-layer prune of sgl-project/DeepSeek-V4-Flash-FP8.
"DeepSeek-V4-Flash-FP8-4layer": "Pinaster",
}
_MEGATRON_MODEL_TYPE = {
"DeepSeek-V4-Flash-FP8": "deepseek-v4-flash",
"DeepSeek-V4-Flash-FP8-4layer": "deepseek-v4-flash-4layer",
}
@dataclass
class ScriptArgs(U.ExecuteTrainConfig):
mode: Literal["normal", "debug_minimal"] = "debug_minimal"
run_id: str = U.create_run_id()
model_org: str = ""
model_name: Literal[
"DeepSeek-V4-Flash-FP8",
"DeepSeek-V4-Flash-FP8-4layer",
] = "DeepSeek-V4-Flash-FP8"
task: Literal["dapo_aime", "gsm8k"] = "dapo_aime"
enable_eval: bool = True
hf_checkpoint: str | None = None
data_dir: str = "/root/datasets"
model_dir: str = "/root/models"
# Defaults to model_dir. Set explicitly when shared NFS -> per-node local NVMe copy is needed.
model_local_dir: str | None = None
save_dir: str = "/root/models"
megatron_path: str = "/root/Megatron-LM"
# performance configs
num_gpus_per_node: int = 8
# use colocate by default. will switch to disaggregated mode when 0 < rollout_num_nodes < num_nodes
rollout_num_nodes: int = 0
colocate: bool = field(init=False)
actor_num_nodes: int = field(init=False)
actor_num_gpus_per_node: int = field(init=False)
rollout_num_gpus: int = field(init=False)
optimizer_offload: bool = True
use_fault_tolerance: bool = True
# debug configs
dump_details: bool = False
debug_train_run_id: str | None = None
debug_train_rollout_id: str | None = None
debug_data_root: str = "/root/shared_data"
skip_saving: bool = False
# precision configs
enable_r3: bool = True
train_deterministic: bool = True
# Megatron-side training precision: blockwise FP8 128x128 GEMMs (fp32 scales) when True,
# BF16 when False. Rollout always serves the source FP8 checkpoint either way.
fp8_training: bool = True
enable_mis: bool = False
# pass any extra sglang/miles/megatron args through `--extra-args '--your-arg'`
extra_args: str = ""
def __post_init__(self):
if not self.model_org:
self.model_org = _DEFAULT_MODEL_ORG[self.model_name]
if self.model_local_dir is None:
self.model_local_dir = self.model_dir
assert self.rollout_num_nodes >= 0
assert self.rollout_num_nodes < self.num_nodes
self.colocate = self.rollout_num_nodes == 0
self.actor_num_nodes = self.num_nodes - self.rollout_num_nodes
self.actor_num_gpus_per_node = self.num_gpus_per_node
if self.colocate:
self.rollout_num_gpus = self.num_nodes * self.num_gpus_per_node
else:
self.rollout_num_gpus = self.rollout_num_nodes * self.num_gpus_per_node
@property
def megatron_model_type(self):
return _MEGATRON_MODEL_TYPE[self.model_name]
@property
def torch_dist_name(self):
return f"{self.model_name}_torch_dist"
@property
def bf16_name(self):
return f"{self.model_name}-bf16"
def _download_dataset(args: ScriptArgs):
"""Download the task-specific dataset(s)."""
match args.task:
case "dapo_aime":
U.hf_download_dataset("zhuzilin/dapo-math-17k", data_dir=args.data_dir)
U.hf_download_dataset("zhuzilin/aime-2024", data_dir=args.data_dir)
case "gsm8k":
U.hf_download_dataset("zhuzilin/gsm8k", data_dir=args.data_dir)
def _hf_checkpoint_path(args: ScriptArgs) -> str:
"""Resolve hf_checkpoint path: explicit override wins, else {model_dir}/{model_name}."""
return args.hf_checkpoint or f"{args.model_dir}/{args.model_name}"
def _ensure_4layer_model_type(args: ScriptArgs):
"""Undo the old deepseek_ref workaround for local 4-layer prunes."""
if args.model_name != "DeepSeek-V4-Flash-FP8-4layer":
return
cfg = Path(_hf_checkpoint_path(args)) / "config.json"
if not cfg.exists():
return
text = cfg.read_text()
if '"model_type": "deepseek_ref"' in text:
cfg.write_text(text.replace('"model_type": "deepseek_ref"', '"model_type": "deepseek_v4"'))
print(f"[patch] {cfg}: model_type deepseek_ref -> deepseek_v4")
def _prepare_download(args: ScriptArgs):
"""Download HF checkpoint + task dataset. Idempotent: hf skips existing blobs."""
U.exec_command(f"mkdir -p {args.model_dir} {args.data_dir}")
# Only download if the user has NOT supplied a pre-existing checkpoint dir.
# (prepare_single / train with --hf-checkpoint bypass this.)
if args.hf_checkpoint is None:
dest = f"{args.model_dir}/{args.model_name}"
U.exec_command(f"hf download {args.model_org}/{args.model_name} " f"--local-dir {dest}")
_ensure_4layer_model_type(args)
_download_dataset(args)
@app.command()
@U.dataclass_cli
def prepare_download(args: ScriptArgs):
"""Download HF checkpoint + dataset from HuggingFace. Run on one node (shared NFS)."""
_prepare_download(args)
def _prepare_single(args: ScriptArgs):
_download_dataset(args)
src = _hf_checkpoint_path(args)
U.fp8_cast_bf16(
path_src=src,
path_dst=f"{args.model_dir}/{args.bf16_name}/",
)
@app.command()
@U.dataclass_cli
def prepare_single(args: ScriptArgs):
"""FP8 -> BF16 cast for Megatron. Needs --hf-checkpoint (or pre-downloaded). One node."""
_prepare_single(args)
def _prepare_spmd(args: ScriptArgs):
is_4layer = args.model_name == "DeepSeek-V4-Flash-FP8-4layer"
actor_num_nodes = args.actor_num_nodes
actor_num_gpus_per_node = args.actor_num_gpus_per_node
extra_args = "--expert-tensor-parallel-size 1 --context-parallel-size 1 "
if actor_num_nodes == 1 and is_4layer:
extra_args += (
"--tensor-model-parallel-size 1 " "--pipeline-model-parallel-size 1 " "--expert-model-parallel-size 1 "
)
elif actor_num_nodes == 1 and args.model_name == "DeepSeek-V4-Flash-FP8":
extra_args += (
"--tensor-model-parallel-size 1 " "--pipeline-model-parallel-size 1 " "--expert-model-parallel-size 8 "
)
else:
raise NotImplementedError(
f"No verified SPMD conversion config for {args.model_name} "
f"({actor_num_nodes} actor nodes x {actor_num_gpus_per_node} GPUs/node). "
f"Please specify your conversion parallel config in `run_deepseek_v4.py`."
)
num_gpus_for_convert = actor_num_gpus_per_node
if is_4layer:
num_gpus_for_convert = min(num_gpus_for_convert, 4)
U.convert_checkpoint(
model_name=args.model_name,
hf_checkpoint=f"{args.model_dir}/{args.bf16_name}",
megatron_model_type=args.megatron_model_type,
num_gpus_per_node=num_gpus_for_convert,
multinode=True if actor_num_nodes > 1 else False,
num_nodes=actor_num_nodes,
extra_args=extra_args,
dir_dst=f"{args.model_dir}",
megatron_path=args.megatron_path,
)
@app.command()
@U.dataclass_cli
def prepare_spmd(args: ScriptArgs):
_prepare_spmd(args)
@app.command()
@U.dataclass_cli
def prepare_cp(args: ScriptArgs):
_prepare_cp(args)
def _prepare_cp(args: ScriptArgs):
U.rsync_simple(
path_src=f"{args.model_dir}/{args.torch_dist_name}",
path_dst=f"{args.model_local_dir}/{args.torch_dist_name}",
num_nodes=args.num_nodes,
)
U.rsync_simple(
path_src=f"{args.model_dir}/{args.model_name}",
path_dst=f"{args.model_local_dir}/{args.model_name}",
num_nodes=args.num_nodes,
)
def _get_parallel_config(args: ScriptArgs) -> str:
"""Return parallel config args for tested GPU configurations.
Only includes configurations that have been verified to work.
Raises NotImplementedError for untested configurations.
"""
actor_num_nodes = args.actor_num_nodes
actor_num_gpus_per_node = args.actor_num_gpus_per_node
total_gpus = actor_num_nodes * actor_num_gpus_per_node
# Single-node smoke-test configs
if actor_num_nodes == 1:
return (
f"--tensor-model-parallel-size {actor_num_gpus_per_node} "
"--sequence-parallel "
"--pipeline-model-parallel-size 1 "
"--context-parallel-size 1 "
f"--expert-model-parallel-size {actor_num_gpus_per_node} "
"--expert-tensor-parallel-size 1 "
)
if actor_num_gpus_per_node == 8:
if total_gpus == 32: # 4 nodes x 8 GPUs (MI355X, full Flash): TP8/PP4/EP8, 43 layers = 11+11+11+10
return (
"--tensor-model-parallel-size 8 "
"--sequence-parallel "
"--pipeline-model-parallel-size 4 "
"--decoder-first-pipeline-num-layers 11 "
"--decoder-last-pipeline-num-layers 10 "
"--context-parallel-size 1 "
"--expert-model-parallel-size 8 "
"--expert-tensor-parallel-size 1 "
)
raise NotImplementedError(
f"No pre-set parallel config for {total_gpus} GPUs. "
f"Please specify your parallel config in `run_deepseek_v4._get_parallel_config`."
)
def _train(args: ScriptArgs):
print(f"[precision] fp8_training={args.fp8_training}")
print(
f"running on {args.num_nodes} nodes "
f"({args.actor_num_nodes} actor nodes x {args.actor_num_gpus_per_node} GPUs/node, "
f"{args.rollout_num_gpus} rollout GPUs, colocate={args.colocate})"
)
_ensure_4layer_model_type(args)
load_save_path = f"{args.save_dir}/{args.run_id}/checkpoints"
ckpt_args = f"--hf-checkpoint {args.hf_checkpoint} " f"--ref-load {args.model_local_dir}/{args.torch_dist_name} "
if not args.skip_saving:
ckpt_args += (
f"--load {load_save_path} " f"--save {load_save_path} " "--save-interval 20 " "--save-retain-interval 20 "
)
rollout_args = (
"--label-key label "
"--apply-chat-template "
"--rollout-shuffle "
"--rm-type math "
"--num-rollout 3000 "
"--rollout-batch-size 32 "
"--n-samples-per-prompt 8 "
"--rollout-temperature 0.8 "
"--num-steps-per-rollout 1 "
"--balance-data "
)
if args.mode != "debug_minimal":
rollout_args += (
"--over-sampling-batch-size 512 "
"--dynamic-sampling-filter-path miles.rollout.filter_hub.dynamic_sampling_filters.check_reward_nonzero_std "
)
eval_args = ""
if args.enable_eval:
eval_args += "--eval-interval 20 " "--eval-top-p 0.7 "
match args.task:
case "dapo_aime":
rollout_args += (
f"--prompt-data {args.data_dir}/dapo-math-17k/dapo-math-17k.jsonl "
"--input-key prompt "
f"--rollout-max-response-len 4096 "
"""--apply-chat-template-kwargs '{"thinking_mode":"thinking"}' """
)
eval_args += (
f"--eval-prompt-data aime {args.data_dir}/aime-2024/aime-2024.jsonl "
"--n-samples-per-eval-prompt 8 "
"--eval-max-response-len 4096 "
)
case "gsm8k":
rollout_args += (
f"--prompt-data {args.data_dir}/gsm8k/train.parquet "
"--input-key messages "
"--rollout-max-response-len 256 "
)
eval_args += (
f"--eval-prompt-data gsm8k {args.data_dir}/gsm8k/test.parquet "
"--n-samples-per-eval-prompt 1 "
"--eval-max-response-len 256 "
)
perf_args = _get_parallel_config(args)
perf_args += (
"--recompute-granularity full "
"--recompute-method uniform "
"--recompute-num-layers 1 "
"--micro-batch-size 1 "
"--max-tokens-per-gpu 2048 "
)
grpo_args = (
"--advantage-estimator grpo "
"--kl-loss-coef 0.00 "
"--kl-loss-type low_var_kl "
"--entropy-coef 0.00 "
"--eps-clip 0.2 "
"--eps-clip-high 0.28 "
)
optimizer_args = (
"--optimizer adam "
"--lr 1e-6 "
"--lr-decay-style constant "
"--weight-decay 0.1 "
"--adam-beta1 0.9 "
"--adam-beta2 0.98 "
)
if args.optimizer_offload:
optimizer_args += (
"--optimizer-cpu-offload " "--use-precision-aware-optimizer " "--overlap-cpu-optimizer-d2h-h2d "
)
if args.actor_num_nodes == 4:
# 4-node PP4 memory balance: partial optimizer offload (keep ~25% on GPU) + keep train
# weights on GPU; pair with --sglang-mem-fraction-static 0.6.
optimizer_args += "--optimizer-offload-fraction 0.75 " "--no-offload-train "
sglang_world_size = 4
sglang_tp_size = 4
sglang_dp_size = 1
sglang_ep_size = 4
sglang_args = (
f"--rollout-num-gpus-per-engine {sglang_world_size} "
f"--sglang-tp-size {sglang_tp_size} "
f"--sglang-dp-size {sglang_dp_size} "
f"--sglang-ep-size {sglang_ep_size} "
"--router-health-success-threshold 1 "
"--router-health-check-interval-secs 15 "
"--router-health-failure-threshold 40 " # TODO improve
# gfx950: DSv4 sgl-kernel topk_v2 is CUDA-only -> route DSA topk through torch + disable cuda-graph.
"--sglang-disable-cuda-graph "
"--sglang-dsa-topk-backend torch "
)
extra_env_vars = {
"SGLANG_SKIP_CHECKPOINT_LOAD_CHECK": "1",
"SGLANG_DSV4_FP4_EXPERTS": "0",
"SGLANG_HEALTH_CHECK_TIMEOUT": "120",
"SGLANG_DG_CACHE_DIR_PER_PROCESS": "1",
"SGLANG_OPT_FP8_WO_A_GEMM": "0",
# ROCm/gfx950 rollout kernel knobs
"SGLANG_HACK_FLASHMLA_BACKEND": "triton",
"SGLANG_FP8_PAGED_MQA_LOGITS_TORCH": "1",
"SGLANG_DSA_TOPK_BROADCAST": "1",
"SGLANG_OPT_USE_TILELANG_INDEXER": "true",
"SGLANG_OPT_USE_AITER_INDEXER": "false",
"SGLANG_OPT_USE_TILELANG_MHC_PRE": "false",
"SGLANG_OPT_USE_TILELANG_MHC_POST": "false",
"SGLANG_OPT_DEEPGEMM_HC_PRENORM": "false",
"SGLANG_OPT_USE_FUSED_COMPRESS": "true",
"SGLANG_OPT_USE_FUSED_COMPRESS_TRITON": "true",
"SGLANG_OPT_USE_JIT_INDEXER_METADATA": "false",
"SGLANG_OPT_USE_TOPK_V2": "false",
"SGLANG_OPT_USE_COMPRESSOR_V2": "false",
"SGLANG_OPT_USE_MULTI_STREAM_OVERLAP": "false",
"SGLANG_ROCM_USE_MULTI_STREAM": "false",
"SGLANG_OPT_USE_CUSTOM_ALL_REDUCE_V2": "0",
"AITER_BF16_FP8_MOE_BOUND": "0",
}
misc_args = (
"--attention-dropout 0.0 "
"--hidden-dropout 0.0 "
"--attention-softmax-in-fp32 "
f"--update-weight-buffer-size {1 * 1024 ** 3} "
f"--actor-num-nodes {args.actor_num_nodes} "
f"--actor-num-gpus-per-node {args.actor_num_gpus_per_node} "
f"--num-gpus-per-node {args.num_gpus_per_node} "
"--train-memory-margin-bytes 3221225472 "
"--sglang-mem-fraction-static 0.7 "
"--sglang-watchdog-timeout 1800 " # ROCm: slow aiter gemm tune under colocate; avoid watchdog SIGQUIT
"--accumulate-allreduce-grads-in-fp32 "
"--model-name deepseekv4 " # for mbridge load
"--qkv-format bshd "
"--moe-router-freeze-gate "
"--freeze-e-score-correction-bias "
"--rollout-health-check-interval 300 "
"--rollout-health-check-timeout 300 "
)
if args.colocate:
misc_args += "--colocate "
else:
misc_args += f"--rollout-num-gpus {args.rollout_num_gpus} "
if args.dump_details:
misc_args += f"--dump-details {args.debug_data_root}/{args.run_id}/dump_details "
if args.enable_mis:
misc_args += (
"--use-tis "
"--custom-config-path examples/train_infer_mismatch_helper/mis.yaml "
"--custom-tis-function-path examples.train_infer_mismatch_helper.mis.compute_mis_weights_with_cp "
)
if args.use_fault_tolerance:
misc_args += "--use-fault-tolerance "
if args.debug_train_run_id is not None:
if args.debug_train_rollout_id is None:
args.debug_train_rollout_id = 1
misc_args += (
f"--load-debug-rollout-data "
f"{args.debug_data_root}/{args.debug_train_run_id}/dump_details/rollout_data/{args.debug_train_rollout_id}.pt "
)
misc_args += "--debug-train-only "
if args.enable_r3:
misc_args += "--use-rollout-routing-replay "
# Skip indexer-replay for now
# misc_args += "--use-rollout-indexer-replay "
# Route replay through the miles python router: the Rust router drops return_routed_experts
# on /generate passthrough, so routed_experts never reaches the scheduler.
misc_args += "--use-miles-router "
if args.train_deterministic:
misc_args += "--deterministic-mode "
extra_env_vars |= {
"NCCL_ALGO": "Ring",
"NVTE_ALLOW_NONDETERMINISTIC_ALGO": "0",
"CUBLAS_WORKSPACE_CONFIG": ":4096:8",
}
if args.fp8_training:
misc_args += "--transformer-impl transformer_engine " "--bf16 " "--fp8-format e4m3 " "--fp8-recipe blockwise "
# gfx950 uses blockwise FP8 with fp32 scales.
misc_args += """--train-env-vars '{"NVTE_FP8_BLOCK_SCALING_FP32_SCALES":"1"}' """
# ROCm TE MoE FP8 lacks fused wgrad accumulation; disable the fusion.
misc_args += "--no-gradient-accumulation-fusion "
train_args = (
f"{ckpt_args} "
f"{rollout_args} "
f"{optimizer_args} "
f"{grpo_args} "
f"{U.get_default_wandb_args(__file__, run_id=args.run_id)} "
f"{perf_args} "
f"{eval_args} "
f"{sglang_args} "
f"{misc_args} "
f"{args.extra_args} "
)
U.execute_train(
train_args=train_args,
config=args,
num_gpus_per_node=args.num_gpus_per_node,
megatron_model_type=args.megatron_model_type,
extra_env_vars={**extra_env_vars},
megatron_path=args.megatron_path,
)
@app.command()
@U.dataclass_cli
def train(args: ScriptArgs):
"""Run training. Assumes data/model/torch_dist are already prepared on {model_local_dir}."""
_train(args)
@app.command()
@U.dataclass_cli
def full_train(args: ScriptArgs):
_prepare_download(args)
bf16_dir = Path(f"{args.model_dir}/{args.bf16_name}")
bf16_sentinel = bf16_dir / "model.safetensors.index.json"
if not bf16_sentinel.exists():
_prepare_single(args)
else:
print(f"[full_train] Skipping FP8->BF16 cast: {bf16_sentinel} already exists.")
torch_dist_dir = Path(f"{args.model_dir}/{args.torch_dist_name}")
torch_dist_sentinel = torch_dist_dir / "latest_checkpointed_iteration.txt"
if not torch_dist_sentinel.exists():
_prepare_spmd(args)
else:
print(f"[full_train] Skipping BF16->torch_dist conversion: {torch_dist_sentinel} already exists.")
if args.model_local_dir != args.model_dir:
_prepare_cp(args)
else:
print(f"[full_train] Skipping rsync: model_local_dir == model_dir ({args.model_dir})")
if args.hf_checkpoint is None:
args.hf_checkpoint = f"{args.model_local_dir}/{args.model_name}"
_train(args)
if __name__ == "__main__":
app()