[AMD][LoRA] Add a Qwen3-4B LoRA RL launcher (#3183)

This commit is contained in:
Yikai Zhang
2026-09-14 14:33:01 -07:00
committed by GitHub
parent 9a9c3eb386
commit 664d332796
6 changed files with 469 additions and 1 deletions
+1 -1
View File
@@ -68,7 +68,7 @@ full-scale experiment evidence. It is not an exhaustive model whitelist.
| Implementation | Model family | Architecture exercised | Evidence | Notes |
|---|---|---|---|---|
| Bridge | Qwen2.5 0.5B / 3B | Dense | [0.5B CUDA and ROCm E2E](https://github.com/radixark/miles/blob/main/tests/e2e/lora/test_lora_qwen2.5_0.5B.py), [3B disaggregated recipe](https://github.com/radixark/miles/blob/main/examples/lora/run-qwen2.5-3B-megatron-lora-disaggregated.sh) | The simplest starting point; `all-linear` works. |
| Bridge | Qwen3 4B | Dense | [Single-LoRA recipe](https://github.com/radixark/miles/blob/main/examples/lora/run-qwen3-4B-megatron-lora.sh), [multi-LoRA recipe](https://github.com/radixark/miles/tree/main/examples/multi_lora) | Used by the current multi-adapter example. |
| Bridge | Qwen3 4B | Dense | [Single-LoRA recipe](https://github.com/radixark/miles/blob/main/examples/lora/run-qwen3-4B-megatron-lora.sh), [AMD launcher](https://github.com/radixark/miles/blob/main/scripts/amd/run_qwen3_4b_lora.py), [multi-LoRA recipe](https://github.com/radixark/miles/tree/main/examples/multi_lora) | The AMD launcher uses Triton attention and LoRA backends. |
| Bridge | GPT-OSS 20B | MoE | [Recipe](https://github.com/radixark/miles/blob/main/examples/lora/run-gpt-oss-20B-megatron-moe-lora.sh), [MoE LoRA E2E](https://github.com/radixark/miles/blob/main/tests/e2e/megatron/model_scripts/test_gpt_oss_20b_moe_lora_ci.py) | Uses the SGLang `triton` LoRA backend. |
| Bridge | Kimi K2.5 | Multimodal MoE + MLA | [16-node recipe](https://github.com/radixark/miles/blob/main/examples/lora/run-kimi-k25-megatron-lora.sh) | Demonstrates shared-outer expert LoRA and an INT4 rollout / fake-QAT setup. |
| Bridge | GLM-5 / 5.1 / 5.2 744B-A40B | MoE + MLA + DSA | [GLM-5.1 launcher](https://github.com/radixark/miles/blob/main/scripts/run_glm5_1_744b_a40b_lora.py), [GLM-5.2 launcher](https://github.com/radixark/miles/blob/main/scripts/run_glm5_2_744b_a40b_lora.py) | CI covers reduced 6-layer / 5-layer checkpoints; historical full-744B results are described below. |
+250
View File
@@ -0,0 +1,250 @@
"""Qwen3-4B GRPO LoRA training script for AMD (MI350X / MI355X).
This is the AMD counterpart of
``examples/lora/run-qwen3-4b-megatron-lora-result.sh``. It keeps the validated
Megatron-Bridge LoRA recipe and pins SGLang's attention and LoRA backends to Triton,
because the CUDA-only rollout backends are unavailable on ROCm.
The checkpoint is loaded directly from Hugging Face through Megatron-Bridge; no converted
Megatron checkpoint is required. ``full-train`` downloads the checkpoint and datasets
before submitting the training job.
Args:
--hardware: MI350X or MI355X, which fixes the default GPU count per node.
--num-gpus-per-node: Override the GPU count when only some devices are visible.
--wandb-team: W&B entity (personal account or team); needed if the API key has no
default entity.
--model-dir / --data-dir / --output-dir: Model, dataset, and run output directories.
--enable-eval: Evaluate AIME 2024 every 20 optimizer steps (default: on).
Examples:
python scripts/amd/run_qwen3_4b_lora.py prepare
python scripts/amd/run_qwen3_4b_lora.py train --hardware MI355X --wandb-team my-entity
python scripts/amd/run_qwen3_4b_lora.py full-train --hardware MI355X --wandb-team my-entity
"""
import os
import shlex
from dataclasses import dataclass
from typing import Literal
import typer
import miles.utils.external_utils.command_utils as U
app = typer.Typer()
_HF_REPO = "Qwen/Qwen3-4B"
@dataclass
class ScriptArgs(U.ExecuteTrainConfig):
run_id: str = U.create_run_id()
model_name: str = "Qwen3-4B"
megatron_model_type: str = "qwen3-4B"
hardware: Literal["auto", "MI350X", "MI355X"] = "auto"
num_gpus_per_node: int | None = None
model_dir: str = "/root/models"
data_dir: str = "/root/datasets"
megatron_path: str = "/root/Megatron-LM"
save_interval: int = 50
# LoRA: inherited from the validated Qwen3-4B Megatron LoRA recipe.
lora_rank: int = 64
lora_alpha: int = 32
lora_dropout: float = 0.0
target_modules: str = "all-linear"
lora_base_cpu_backup: bool = True
# Rollout and optimizer shape: 8 prompts x 8 samples = one 64-sample update.
num_rollout: int = 100
rollout_batch_size: int = 8
n_samples_per_prompt: int = 8
rollout_max_response_len: int = 8192
global_batch_size: int = 64
lr: float = 2e-5
# One 4B rollout engine per GPU.
rollout_num_gpus_per_engine: int = 1
sglang_mem_fraction_static: float = 0.4
enable_eval: bool = True
enable_wandb: bool = True
wandb_team: str | None = None
extra_args: str = ""
def _set_rocm_environment() -> None:
"""Preserve the caller's ROCm device selection through Ray startup."""
os.environ.setdefault("RAY_EXPERIMENTAL_NOSET_HIP_VISIBLE_DEVICES", "1")
os.environ.setdefault("RAY_EXPERIMENTAL_NOSET_CUDA_VISIBLE_DEVICES", "1")
if hip_visible_devices := os.environ.get("HIP_VISIBLE_DEVICES"):
os.environ["CUDA_VISIBLE_DEVICES"] = hip_visible_devices
# Avoid execute_train's NVIDIA topology probe and disable unsupported NVLink SHARP.
os.environ.setdefault("NCCL_NVLS_ENABLE", "0")
def _resolve_num_gpus(args: ScriptArgs) -> tuple[str, int]:
hardware = U.resolve_hardware(args)
num_gpus = args.num_gpus_per_node or U.NUM_GPUS_OF_HARDWARE[hardware]
return hardware, num_gpus
def _download_inputs(args: ScriptArgs) -> None:
U.exec_command_cpu(f"mkdir -p {args.model_dir} {args.data_dir}")
U.exec_command_cpu(f"hf download {_HF_REPO} --local-dir {args.model_dir}/{args.model_name}")
U.hf_download_dataset("zhuzilin/dapo-math-17k", data_dir=args.data_dir)
if args.enable_eval:
U.hf_download_dataset("zhuzilin/aime-2024", data_dir=args.data_dir)
def _get_wandb_args(args: ScriptArgs) -> str:
if not args.enable_wandb:
return ""
wandb_args = U.get_default_wandb_args(__file__, run_id=args.run_id)
if wandb_args and args.wandb_team:
wandb_args += f"--wandb-team {shlex.quote(args.wandb_team)} "
return wandb_args
def _execute(args: ScriptArgs) -> None:
_set_rocm_environment()
hardware, num_gpus = _resolve_num_gpus(args)
print(f"[run] Qwen3-4B LoRA on {hardware}: {num_gpus} GPUs, rollout TP=1")
checkpoint_dir = f"{args.output_dir}/{args.run_id}/checkpoints"
ckpt_args = (
f"--hf-checkpoint {args.model_dir}/{args.model_name} "
"--megatron-to-hf-mode bridge "
f"--save {checkpoint_dir} "
f"--save-interval {args.save_interval} "
)
lora_args = (
f"--lora-rank {args.lora_rank} "
f"--lora-alpha {args.lora_alpha} "
f"--lora-dropout {args.lora_dropout} "
f'--target-modules "{args.target_modules}" '
)
if args.lora_base_cpu_backup:
lora_args += "--lora-base-cpu-backup "
rollout_args = (
f"--prompt-data {args.data_dir}/dapo-math-17k/dapo-math-17k.jsonl "
"--input-key prompt "
"--label-key label "
"--apply-chat-template "
"--rollout-shuffle "
"--balance-data "
"--rm-type deepscaler "
f"--num-rollout {args.num_rollout} "
f"--rollout-batch-size {args.rollout_batch_size} "
f"--n-samples-per-prompt {args.n_samples_per_prompt} "
f"--rollout-max-response-len {args.rollout_max_response_len} "
"--rollout-temperature 1 "
f"--global-batch-size {args.global_batch_size} "
)
eval_args = ""
if args.enable_eval:
eval_args = (
"--eval-interval 20 "
f"--eval-prompt-data aime24 {args.data_dir}/aime-2024/aime-2024.jsonl "
"--n-samples-per-eval-prompt 16 "
"--eval-max-response-len 16384 "
"--eval-top-p 1 "
)
optimizer_args = (
"--optimizer adam "
f"--lr {args.lr} "
"--lr-decay-style constant "
"--weight-decay 0.1 "
"--adam-beta1 0.9 "
"--adam-beta2 0.98 "
)
grpo_args = (
"--advantage-estimator grpo "
"--kl-loss-coef 0.00 "
"--kl-loss-type low_var_kl "
"--kl-coef 0.00 "
"--entropy-coef 0.00 "
"--eps-clip 0.2 "
"--eps-clip-high 0.28 "
)
sglang_args = (
f"--rollout-num-gpus-per-engine {args.rollout_num_gpus_per_engine} "
f"--sglang-mem-fraction-static {args.sglang_mem_fraction_static} "
"--sglang-decode-log-interval 1000 "
"--sglang-chunked-prefill-size 4096 "
"--sglang-lora-backend triton "
"--sglang-attention-backend triton "
)
misc_args = (
"--train-backend megatron "
"--attention-dropout 0.0 "
"--hidden-dropout 0.0 "
"--accumulate-allreduce-grads-in-fp32 "
"--attention-softmax-in-fp32 "
"--calculate-per-token-loss "
f"--actor-num-nodes {args.num_nodes} "
f"--actor-num-gpus-per-node {num_gpus} "
f"--num-gpus-per-node {num_gpus} "
"--colocate "
)
train_args = (
f"{ckpt_args} "
f"{lora_args} "
f"{rollout_args} "
f"{eval_args} "
f"{optimizer_args} "
f"{grpo_args} "
f"{_get_wandb_args(args)} "
f"{sglang_args} "
f"{misc_args} "
f"{args.extra_args} "
)
U.execute_train(
train_args=train_args,
config=args,
num_gpus_per_node=num_gpus,
megatron_model_type=args.megatron_model_type,
megatron_path=args.megatron_path,
)
@app.command()
@U.dataclass_cli
def prepare(args: ScriptArgs) -> None:
"""Download Qwen3-4B and the configured training/evaluation datasets."""
_download_inputs(args)
@app.command()
@U.dataclass_cli
def train(args: ScriptArgs) -> None:
"""Run GRPO LoRA training using inputs already present on the node."""
_execute(args)
@app.command()
@U.dataclass_cli
def full_train(args: ScriptArgs) -> None:
"""Download the inputs and run GRPO LoRA training."""
_download_inputs(args)
_execute(args)
@app.callback()
def _callback() -> None:
pass
if __name__ == "__main__":
app()
@@ -70,6 +70,7 @@ _SCRIPTS_WHOSE_DEFAULTS_ARE_UNSUPPORTED: dict[str, Callable[[Path], dict[str, ob
_HARDWARE_A_RECORDING_REPRESENTS = {
"scripts/amd/run_qwen3_30b_a3b.py": "MI355X",
"scripts/amd/run_qwen3_4b.py": "MI355X",
"scripts/amd/run_qwen3_4b_lora.py": "MI355X",
"scripts/run_deepseek_v32.py": "B200",
"scripts/run_glm45_355b_a32b.py": "GB200",
"scripts/run_joy_ai_llm_flash.py": "B200",
@@ -0,0 +1,109 @@
### 0
mkdir -p /root/models /root/datasets
### 1
hf download Qwen/Qwen3-4B
--local-dir /root/models/Qwen3-4B
### 2
hf download
--repo-type dataset zhuzilin/dapo-math-17k
--local-dir /root/datasets/dapo-math-17k
### 3
hf download
--repo-type dataset zhuzilin/aime-2024
--local-dir /root/datasets/aime-2024
### 4
pkill -9 sglang; sleep 3; ray stop
--force; pkill -9 ray; pkill -9 miles; sleep 3; pkill -9 ray; pkill -9 miles; pkill -9 redis; true;
### 5
export PYTHONUNBUFFERED=1 && ray start
--head
--node-ip-address 127.0.0.1
--num-gpus 8
--disable-usage-stats
### 6
export no_proxy=127.0.0.1 && export PYTHONUNBUFFERED=1 && ray job submit
--address="http://127.0.0.1:8265"
--runtime-env-json='{"env_vars": {"PYTHONUNBUFFERED": "1", "CUDA_DEVICE_MAX_CONNECTIONS": "1", "NCCL_NVLS_ENABLE": "0", "no_proxy": "127.0.0.1,127.0.0.1", "MASTER_ADDR": "127.0.0.1", "PYTHONPATH": "<REPO_ROOT>:/root/Megatron-LM:/frozen/pythonpath"}}'
-- python3 <REPO_ROOT>/train.py
--swiglu
--num-layers 36
--hidden-size 2560
--ffn-hidden-size 9728
--num-attention-heads 32
--group-query-attention
--num-query-groups 8
--use-rotary-position-embeddings
--disable-bias-linear
--normalization RMSNorm
--norm-epsilon 1e-6
--rotary-base 1000000
--vocab-size 151936
--kv-channels 128
--qk-layernorm
--hf-checkpoint /root/models/Qwen3-4B
--megatron-to-hf-mode bridge
--save /root/shared_data/260101-000000-000/checkpoints
--save-interval 50
--lora-rank 64
--lora-alpha 32
--lora-dropout 0.0
--target-modules "all-linear"
--lora-base-cpu-backup
--prompt-data /root/datasets/dapo-math-17k/dapo-math-17k.jsonl
--input-key prompt
--label-key label
--apply-chat-template
--rollout-shuffle
--balance-data
--rm-type deepscaler
--num-rollout 100
--rollout-batch-size 8
--n-samples-per-prompt 8
--rollout-max-response-len 8192
--rollout-temperature 1
--global-batch-size 64
--eval-interval 20
--eval-prompt-data aime24 /root/datasets/aime-2024/aime-2024.jsonl
--n-samples-per-eval-prompt 16
--eval-max-response-len 16384
--eval-top-p 1
--optimizer adam
--lr 2e-05
--lr-decay-style constant
--weight-decay 0.1
--adam-beta1 0.9
--adam-beta2 0.98
--advantage-estimator grpo
--kl-loss-coef 0.00
--kl-loss-type low_var_kl
--kl-coef 0.00
--entropy-coef 0.00
--eps-clip 0.2
--eps-clip-high 0.28
--use-wandb
--wandb-project miles-run_qwen3_4b_lora
--wandb-group 260101-000000-000
--wandb-key 'frozen-wandb-api-key'
--disable-wandb-random-suffix
--rollout-num-gpus-per-engine 1
--sglang-mem-fraction-static 0.4
--sglang-decode-log-interval 1000
--sglang-chunked-prefill-size 4096
--sglang-lora-backend triton
--sglang-attention-backend triton
--train-backend megatron
--attention-dropout 0.0
--hidden-dropout 0.0
--accumulate-allreduce-grads-in-fp32
--attention-softmax-in-fp32
--calculate-per-token-loss
--actor-num-nodes 1
--actor-num-gpus-per-node 8
--num-gpus-per-node 8
--colocate
@@ -0,0 +1,16 @@
### 0
mkdir -p /root/models /root/datasets
### 1
hf download Qwen/Qwen3-4B
--local-dir /root/models/Qwen3-4B
### 2
hf download
--repo-type dataset zhuzilin/dapo-math-17k
--local-dir /root/datasets/dapo-math-17k
### 3
hf download
--repo-type dataset zhuzilin/aime-2024
--local-dir /root/datasets/aime-2024
@@ -0,0 +1,92 @@
### 0
pkill -9 sglang; sleep 3; ray stop
--force; pkill -9 ray; pkill -9 miles; sleep 3; pkill -9 ray; pkill -9 miles; pkill -9 redis; true;
### 1
export PYTHONUNBUFFERED=1 && ray start
--head
--node-ip-address 127.0.0.1
--num-gpus 8
--disable-usage-stats
### 2
export no_proxy=127.0.0.1 && export PYTHONUNBUFFERED=1 && ray job submit
--address="http://127.0.0.1:8265"
--runtime-env-json='{"env_vars": {"PYTHONUNBUFFERED": "1", "CUDA_DEVICE_MAX_CONNECTIONS": "1", "NCCL_NVLS_ENABLE": "0", "no_proxy": "127.0.0.1,127.0.0.1", "MASTER_ADDR": "127.0.0.1", "PYTHONPATH": "<REPO_ROOT>:/root/Megatron-LM:/frozen/pythonpath"}}'
-- python3 <REPO_ROOT>/train.py
--swiglu
--num-layers 36
--hidden-size 2560
--ffn-hidden-size 9728
--num-attention-heads 32
--group-query-attention
--num-query-groups 8
--use-rotary-position-embeddings
--disable-bias-linear
--normalization RMSNorm
--norm-epsilon 1e-6
--rotary-base 1000000
--vocab-size 151936
--kv-channels 128
--qk-layernorm
--hf-checkpoint /root/models/Qwen3-4B
--megatron-to-hf-mode bridge
--save /root/shared_data/260101-000000-000/checkpoints
--save-interval 50
--lora-rank 64
--lora-alpha 32
--lora-dropout 0.0
--target-modules "all-linear"
--lora-base-cpu-backup
--prompt-data /root/datasets/dapo-math-17k/dapo-math-17k.jsonl
--input-key prompt
--label-key label
--apply-chat-template
--rollout-shuffle
--balance-data
--rm-type deepscaler
--num-rollout 100
--rollout-batch-size 8
--n-samples-per-prompt 8
--rollout-max-response-len 8192
--rollout-temperature 1
--global-batch-size 64
--eval-interval 20
--eval-prompt-data aime24 /root/datasets/aime-2024/aime-2024.jsonl
--n-samples-per-eval-prompt 16
--eval-max-response-len 16384
--eval-top-p 1
--optimizer adam
--lr 2e-05
--lr-decay-style constant
--weight-decay 0.1
--adam-beta1 0.9
--adam-beta2 0.98
--advantage-estimator grpo
--kl-loss-coef 0.00
--kl-loss-type low_var_kl
--kl-coef 0.00
--entropy-coef 0.00
--eps-clip 0.2
--eps-clip-high 0.28
--use-wandb
--wandb-project miles-run_qwen3_4b_lora
--wandb-group 260101-000000-000
--wandb-key 'frozen-wandb-api-key'
--disable-wandb-random-suffix
--rollout-num-gpus-per-engine 1
--sglang-mem-fraction-static 0.4
--sglang-decode-log-interval 1000
--sglang-chunked-prefill-size 4096
--sglang-lora-backend triton
--sglang-attention-backend triton
--train-backend megatron
--attention-dropout 0.0
--hidden-dropout 0.0
--accumulate-allreduce-grads-in-fp32
--attention-softmax-in-fp32
--calculate-per-token-loss
--actor-num-nodes 1
--actor-num-gpus-per-node 8
--num-gpus-per-node 8
--colocate