mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
[WIP] feat: lora GLM-5.2 FP8 rollout support (#1678)
Signed-off-by: Yusheng Su <yushengsu.thu@gmail.com>
This commit is contained in:
@@ -32,6 +32,15 @@ Usage:
|
||||
python scripts/run_glm5_2_744b_a40b_lora.py prepare --model-name GLM-5.2_5layer --task dapo-math
|
||||
python scripts/run_glm5_2_744b_a40b_lora.py train --model-name GLM-5.2_5layer --task dapo-math \\
|
||||
--rollout-max-response-len 4096 --num-gpus-per-node 4
|
||||
|
||||
fp8 rollout (train stays bf16; sglang serves <hf_checkpoint>_fp8 via --sglang-config; LoRA
|
||||
adapters still sync per step, only the base-weight sync is skipped). The rollout checkpoint
|
||||
(<hf_checkpoint>_fp8, e.g. the official GLM-5.2 fp8 release) must already exist:
|
||||
python scripts/run_glm5_2_744b_a40b_lora.py train --model-name GLM-5.2_5layer \\
|
||||
--fp8-rollout --num-gpus-per-node 4
|
||||
Expected parity on the toy (gsm8k, 5 rollouts): train_rollout_logprob_abs_diff ~0.27 /
|
||||
train_rollout_kl ~0.058, flat across steps (constant weight-quantization offset), vs
|
||||
~0.010 / ~1.2e-4 with the bf16 rollout.
|
||||
"""
|
||||
|
||||
import os
|
||||
@@ -116,6 +125,8 @@ class ScriptArgs(U.ExecuteTrainConfig):
|
||||
sglang_lora_backend: str = "triton"
|
||||
# serve from a pre-converted _fp8 ckpt (fits engine=8 / 1 node); train stays bf16
|
||||
fp8_rollout: bool = False
|
||||
# rollout-side fp8 checkpoint; defaults to <hf_checkpoint>_fp8 (e.g. GLM-5.2_fp8).
|
||||
fp8_rollout_checkpoint: str | None = None
|
||||
|
||||
enable_wandb: bool = True
|
||||
extra_args: str = ""
|
||||
@@ -123,6 +134,8 @@ class ScriptArgs(U.ExecuteTrainConfig):
|
||||
def __post_init__(self):
|
||||
if self.hf_checkpoint is None:
|
||||
self.hf_checkpoint = f"{self.model_dir}/{self.model_name}"
|
||||
if self.fp8_rollout and self.fp8_rollout_checkpoint is None:
|
||||
self.fp8_rollout_checkpoint = f"{self.hf_checkpoint}_fp8"
|
||||
if self.rollout_max_response_len == 0:
|
||||
self.rollout_max_response_len = 4096 if self.task == "dapo-math" else 512
|
||||
if self.seq_window == 0 and self.task == "dapo-math":
|
||||
@@ -240,10 +253,11 @@ def _train(args: ScriptArgs):
|
||||
if _is_full:
|
||||
# mirrors run_glm5_744b_a40b.py; bf16 ~1488GB needs >=~22 GPUs/engine while fp8
|
||||
# fits engine=min(8, ngpu) on one node
|
||||
_eng = min(8, args.num_gpus_per_node) if args.fp8_rollout else args.rollout_num_gpus_per_engine
|
||||
_decode = "flashmla_kv" if args.fp8_rollout else "flashmla_sparse"
|
||||
_cg = 256 if args.fp8_rollout else 64
|
||||
_kv = "--sglang-kv-cache-dtype fp8_e4m3 " if args.fp8_rollout else ""
|
||||
_fp8_full = args.fp8_rollout and args.model_name == "GLM-5.2"
|
||||
_eng = min(8, args.num_gpus_per_node) if _fp8_full else args.rollout_num_gpus_per_engine
|
||||
_decode = "flashmla_kv" if _fp8_full else "flashmla_sparse"
|
||||
_cg = 256 if _fp8_full else 64
|
||||
_kv = "--sglang-kv-cache-dtype fp8_e4m3 " if _fp8_full else ""
|
||||
sglang_args = (
|
||||
f"--rollout-num-gpus-per-engine {_eng} --sglang-mem-fraction-static {args.sglang_mem_fraction_static} "
|
||||
f"--sglang-enable-dp-attention --sglang-ep-size {_eng} --sglang-dp-size {_eng} "
|
||||
@@ -260,6 +274,23 @@ def _train(args: ScriptArgs):
|
||||
else:
|
||||
sglang_args = f"--rollout-num-gpus-per-engine {args.rollout_num_gpus_per_engine} --sglang-mem-fraction-static {args.sglang_mem_fraction_static} --sglang-cuda-graph-max-bs 64 --sglang-moe-runner-backend triton --sglang-disable-shared-experts-fusion --sglang-lora-backend {args.sglang_lora_backend} --sglang-reasoning-parser glm45 --sglang-tool-call-parser glm47 "
|
||||
|
||||
if args.fp8_rollout:
|
||||
# Serve the fp8 ckpt via --sglang-config; update_weights stays on so the per-step LoRA
|
||||
# sync reaches the engine (the bf16 base sync is already skipped under colocate + backup).
|
||||
sglang_config_path = f"{load_save_path}/sglang_fp8_rollout.yaml"
|
||||
os.makedirs(load_save_path, exist_ok=True)
|
||||
with open(sglang_config_path, "w") as f:
|
||||
f.write(
|
||||
"sglang:\n"
|
||||
" - name: default\n"
|
||||
f" model_path: {args.fp8_rollout_checkpoint}\n"
|
||||
" update_weights: true\n"
|
||||
" server_groups:\n"
|
||||
" - worker_type: regular\n"
|
||||
f" num_gpus: {args.num_gpus_per_node}\n"
|
||||
)
|
||||
sglang_args += f"--sglang-config {sglang_config_path} "
|
||||
|
||||
save_args = f"--save-interval 1 --save {load_save_path} "
|
||||
|
||||
misc_args = f"--attention-dropout 0.0 --hidden-dropout 0.0 --accumulate-allreduce-grads-in-fp32 --attention-softmax-in-fp32 --attention-backend flash --calculate-per-token-loss --use-miles-router --actor-num-nodes 1 --actor-num-gpus-per-node {args.num_gpus_per_node} --num-gpus-per-node {args.num_gpus_per_node} --colocate "
|
||||
|
||||
Reference in New Issue
Block a user