mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
[AMD] Enable DeepSeek-V4-Flash FP8 RL training on MI355X (#1607)
Co-authored-by: Zhiyao Jiang <jessicajiang324@gmail.com>
This commit is contained in:
co-authored by
Zhiyao Jiang
parent
23335971f9
commit
352b4aafcb
+11
-1
@@ -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
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user