[BUG] Is mcore-0.12.0 checkpoint resume correct?
- Dominant language
- Python
- Stars
- 17.9k
- Forks
- 4.5k
- Avg merge
- 4d 3h
- Merged PRs (30d)
- 272
Description
**Describe the bug**
While using chckpoint recume in mcore-0.12, it cannot produce same loss value in the first forward, and the subsequent loss deviation is relatively large. But mcore-0.11 can generate the same loss in the first step, and the subsequent deviation is very small.
**To Reproduce**
test scripts:
```sh
#!/bin/bash
TOKENIZER_MODEL="/workspace/Models/DeepSeek-V3"
export CUDA_DEVICE_MAX_CONNECTIONS=1
MASTER_ADDR=${MASTER_ADDR:-"localhost"}
MASTER_PORT=${MASTER_PORT:-"6000"}
NNODES=${NNODES:-"1"}
NODE_RANK=${RANK:-"0"}
GPUS_PER_NODE=${GPUS_PER_NODE:-"8"}
WORLD_SIZE=$(($GPUS_PER_NODE*$NNODES))
TP=${TP:-"2"}
PP=${PP:-"2"}
EP=${EP:-"2"}
## train params
SEQ_LEN=${SEQ_LEN:-"4096"}
GB=${GB:-"32"} # global batch size
MB=${MB:-"1"} # micro batch size
TRAIN_ITERS=${TRAIN_ITERS:-"10"}
LR_DECAY_ITERS=${LR_DECAY_ITERS:-"7"}
LR_WARMUP_ITERS=${LR_WARMUP_ITERS:-"3"}
EXTRA_ARGS=${EXTRA_ARGS:-""}
# deterministic-mode, cannot use flash-attn, H800 Nan Bug
export NVTE_ALLOW_NONDETERMINISTIC_ALGO=0
export CUBLAS_WORKSPACE_CONFIG=:4096:8
export NCCL_ALGO=Ring
EXTRA_ARGS="$EXTRA_ARGS --deterministic-mode"
## model params
NUM_LAYERS=8
NUM_DENS_LAYERS=1
HIDDEN_SIZE=2048
NUM_ATTN_HEADS=16
INTERMEDIATE_SIZE=10944
MLA_ARGS=(
# --use-precision-aware-optimizer
# --exp-avg-dtype fp16
# --exp-avg-sq-dtype fp16
--no-rope-fusion
--qk-layernorm
# --q-lora-rank 1536
--kv-lora-rank 512
--multi-latent-attention
--qk-head-dim 128
--qk-pos-emb-head-dim 64
--v-head-dim 128
--rotary-scaling-factor 40
)
MOE_ARGS=(
--num-experts 64
--moe-router-topk 6
--moe-ffn-hidden-size 1408
--moe-shared-expert-intermediate-size 2816 # n_shared_experts * ffn_size, 2*1408
--moe-layer-freq "([0]*$NUM_DENS_LAYERS+[1]*($NUM_LAYERS-$NUM_DENS_LAYERS))"
--expert-model-parallel-size $EP
--moe-token-dispatcher-type alltoall
--moe-grouped-gemm
--moe-router-enable-expert-bias
--moe-router-score-function "sigmoid"
# --moe-router-load-balancing-type seq_aux_loss
--moe-router-load-balancing-type aux_loss
--moe-aux-loss-coeff 0.0001
--moe-router-pre-softmax
--moe-router-topk-scaling-factor 6
--moe-per-layer-logging
--moe-router-group-topk 3
--moe-router-num-groups 8
#--attention-backend unfused
# --moe-permute-fusion # must install TE>=V2.1
)
FP8_ARGS=(
# --fp8-format hybrid
# --fp8-margin 0
# --fp8-interval 1
# --fp8-amax-history-len 1024
# --fp8-amax-compute-algo max
# --fp8-param-gather
)
DISTRIBUTED_ARGS=(
--nproc_per_node $GPUS_PER_NODE
--nnodes $NNODES
--node_rank $NODE_RANK
--master_addr $MASTER_ADDR
--master_port $MASTER_PORT
)
MODEL_ARGS=(
--seq-length ${SEQ_LEN}
--num-layers ${NUM_LAYERS}
--hidden-size ${HIDDEN_SIZE}
--ffn-hidden-size ${INTERMEDIATE_SIZE}
--num-attention-heads ${NUM_ATTN_HEADS}
--max-position-embeddings ${SEQ_LEN}
# --add-qkv-bias # qwen2 use qkv bias
--disable-bias-linear
--init-method-std 0.01
--attention-dropout 0.0
--hidden-dropout 0.0
--normalization RMSNorm
--position-embedding-type rope
--swiglu
# FIXME: masked_softmax_fusion?
--no-masked-softmax-fusion
--no-position-embedding
## qwen2 specific
--rotary-base 1000000
--norm-epsilon 1e-6
--untie-embeddings-and-output-weights
)
DATA_ARGS=(
--tokenizer-type HuggingFaceTokenizer
--tokenizer-model ${TOKENIZER_MODEL}
--split 97,2,1
)
if [ -n "${DATA_PATH}" ]; then
DATA_ARGS+=(
--data-path $DATA_PATH
)
else
DATA_ARGS+=(
--mock-data
)
fi
TRAINING_ARGS=(
--micro-batch-size ${MB}
--global-batch-size ${GB}
--lr 1e-4
--train-iters ${TRAIN_ITERS}
--lr-decay-iters ${LR_DECAY_ITERS}
--lr-warmup-iters ${LR_WARMUP_ITERS}
--lr-decay-style cosine
--min-lr 1.0e-5
--weight-decay 0.1
--clip-grad 1.0
--bf16
)
MODEL_PARALLEL_ARGS=(
--tensor-model-parallel-size ${TP}
--pipeline-model-parallel-size ${PP}
--use-distributed-optimizer
# --use-flash-attn
--use-mcore-models
# --overlap-grad-reduce
# --overlap-param-gather
--sequence-parallel
--no-async-tensor-model-parallel-allreduce
# --tp-comm-overlap-rs-dgrad
)
LOGGING_ARGS=(
# --load ./ckpts/1f1b_dsv2_lite_ep2_pp2_dp2 \
--async-save \
--save-interval 5 \
--save ./ckpts/1f1b_dsv2_lite_ep2_pp2_dp2 \
# --no-load-optim \
# --no-load-rng
--eval-interval 1000 \
--eval-iters 0 \
--log-interval 1 \
--log-throughput \
--log-timers-to-tensorboard \
--log-validation-ppl-to-tensorboard \
--log-world-size-to-tensorboard \
--tensorboard-dir "Tensorboards/dual_pipeV_dsv2_lite_ep2_pp2_dp2" \
)
if [ -n "${WANDB_API_KEY}" ]; then
LOGGING_ARGS+=(
--wandb-project ${WANDB_PROJECT:-"test-pretrain"}
--wandb-exp-name ${WANDB_NAME:-"test"}_${MODEL_SIZE}
)
fi
echo "
torchrun ${DISTRIBUTED_ARGS[@]} pretrain_gpt.py \
${MODEL_ARGS[@]} \
${DATA_ARGS[@]} \
${TRAINING_ARGS[@]} \
${MODEL_PARALLEL_ARGS[@]} \
${LOGGING_ARGS[@]} \
${MLA_ARGS[@]} \
${MOE_ARGS[@]} \
${FP8_ARGS[@]} \
${EXTRA_ARGS}
"
NCCL_IB_GID_INDEX=3 NCCL_NET_GDR_LEVEL=SYS torchrun ${DISTRIBUTED_ARGS[@]} pretrain_gpt.py \
${MODEL_ARGS[@]} \
${DATA_ARGS[@]} \
${TRAINING_ARGS[@]} \
${MODEL_PARALLEL_ARGS[@]} \
${LOGGING_ARGS[@]} \
${MLA_ARGS[@]} \
${MOE_ARGS[@]} \
${FP8_ARGS[@]} \
${EXTRA_ARGS}
```
**Expected behavior**
Checkpoint resume can produce same loss.
**Stack trace/logs**
- mcore-0.12.0
```
# mcore-0.12 full 10 steps
[2025-05-09 06:11:17] iteration 1/ 10 | consumed samples: 32 | elapsed time per iteration (ms): 21847.9 | throughput per GPU (TFLOP/s/GPU): 4.5 | learning rate: 3.333333E-05 | global batch size: 32 | lm loss: 1.186396E+01 | load_balancing_loss: 1.000977E+00 | loss scale: 1.0 | grad norm: 6.321 | num zeros: 256053424.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:11:20] iteration 2/ 10 | consumed samples: 64 | elapsed time per iteration (ms): 3286.6 | throughput per GPU (TFLOP/s/GPU): 30.2 | learning rate: 6.666667E-05 | global batch size: 32 | lm loss: 1.186432E+01 | load_balancing_loss: 1.001151E+00 | loss scale: 1.0 | grad norm: 6.350 | num zeros: 256043200.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:11:23] iteration 3/ 10 | consumed samples: 96 | elapsed time per iteration (ms): 3182.4 | throughput per GPU (TFLOP/s/GPU): 31.2 | learning rate: 1.000000E-04 | global batch size: 32 | lm loss: 8.431922E+00 | load_balancing_loss: 1.002058E+00 | loss scale: 1.0 | grad norm: 5.229 | num zeros: 255912208.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:11:27] iteration 4/ 10 | consumed samples: 128 | elapsed time per iteration (ms): 3181.7 | throughput per GPU (TFLOP/s/GPU): 31.2 | learning rate: 8.681980E-05 | global batch size: 32 | lm loss: 5.055252E+00 | load_balancing_loss: 1.000907E+00 | loss scale: 1.0 | grad norm: 4.788 | num zeros: 256074064.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:11:30] iteration 5/ 10 | consumed samples: 160 | elapsed time per iteration (ms): 3178.5 | throughput per GPU (TFLOP/s/GPU): 31.2 | learning rate: 5.500000E-05 | global batch size: 32 | lm loss: 2.602019E+00 | load_balancing_loss: 1.000000E+00 | loss scale: 1.0 | grad norm: 3.338 | num zeros: 255955120.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
saving checkpoint at iteration 5 to ./ckpts/1f1b_dsv2_lite_ep2_pp2_dp2 in torch_dist format
[2025-05-09 06:11:42] iteration 6/ 10 | consumed samples: 192 | elapsed time per iteration (ms): 3395.7 | throughput per GPU (TFLOP/s/GPU): 29.2 | learning rate: 2.318020E-05 | global batch size: 32 | lm loss: 1.016504E+00 | load_balancing_loss: 1.000070E+00 | loss scale: 1.0 | grad norm: 1.714 | num zeros: 255940912.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:11:45] iteration 7/ 10 | consumed samples: 224 | elapsed time per iteration (ms): 3181.8 | throughput per GPU (TFLOP/s/GPU): 31.2 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 6.185102E-01 | load_balancing_loss: 1.000000E+00 | loss scale: 1.0 | grad norm: 0.996 | num zeros: 255883504.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:11:48] iteration 8/ 10 | consumed samples: 256 | elapsed time per iteration (ms): 3183.0 | throughput per GPU (TFLOP/s/GPU): 31.2 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 3.697273E-01 | load_balancing_loss: 1.000000E+00 | loss scale: 1.0 | grad norm: 0.614 | num zeros: 255883488.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:11:51] iteration 9/ 10 | consumed samples: 288 | elapsed time per iteration (ms): 3189.9 | throughput per GPU (TFLOP/s/GPU): 31.1 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 3.007708E-01 | load_balancing_loss: 1.000000E+00 | loss scale: 1.0 | grad norm: 0.518 | num zeros: 256000304.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:11:55] iteration 10/ 10 | consumed samples: 320 | elapsed time per iteration (ms): 3195.1 | throughput per GPU (TFLOP/s/GPU): 31.1 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 2.847902E-01 | load_balancing_loss: 1.000000E+00 | loss scale: 1.0 | grad norm: 0.485 | num zeros: 255930608.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
# mcore-0.12 resume 6-10 steps
[2025-05-09 06:16:38] iteration 6/ 10 | consumed samples: 192 | elapsed time per iteration (ms): 22248.5 | throughput per GPU (TFLOP/s/GPU): 4.5 | learning rate: 2.318020E-05 | global batch size: 32 | lm loss: 1.043662E+00 | load_balancing_loss: 1.000000E+00 | loss scale: 1.0 | grad norm: 2.442 | num zeros: 338313792.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:16:41] iteration 7/ 10 | consumed samples: 224 | elapsed time per iteration (ms): 3220.9 | throughput per GPU (TFLOP/s/GPU): 30.8 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 7.894453E-01 | load_balancing_loss: 1.000279E+00 | loss scale: 1.0 | grad norm: 1.387 | num zeros: 255883472.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:16:44] iteration 8/ 10 | consumed samples: 256 | elapsed time per iteration (ms): 3176.4 | throughput per GPU (TFLOP/s/GPU): 31.3 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 4.729381E-01 | load_balancing_loss: 1.000000E+00 | loss scale: 1.0 | grad norm: 0.788 | num zeros: 255883504.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:16:48] iteration 9/ 10 | consumed samples: 288 | elapsed time per iteration (ms): 3174.3 | throughput per GPU (TFLOP/s/GPU): 31.3 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 3.806268E-01 | load_balancing_loss: 1.000000E+00 | loss scale: 1.0 | grad norm: 0.657 | num zeros: 256000336.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:16:51] iteration 10/ 10 | consumed samples: 320 | elapsed time per iteration (ms): 3213.1 | throughput per GPU (TFLOP/s/GPU): 30.9 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 3.490109E-01 | load_balancing_loss: 1.000000E+00 | loss scale: 1.0 | grad norm: 0.589 | num zeros: 255930624.0 | number of skipped iterations: 0 | number of nan iterations: 0 |
```
- mcore-0.11.0
```
# mcore-0.11 full 10 steps
[2025-05-09 06:22:40] iteration 1/ 10 | consumed samples: 32 | elapsed time per iteration (ms): 21648.3 | throughput per GPU (TFLOP/s/GPU): 17.7 | learning rate: 3.333333E-05 | global batch size: 32 | lm loss: 1.186509E+01 | load_balancing_loss: 8.778212E-01 | loss scale: 1.0 | grad norm: 3.759 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:22:43] iteration 2/ 10 | consumed samples: 64 | elapsed time per iteration (ms): 3156.5 | throughput per GPU (TFLOP/s/GPU): 121.4 | learning rate: 6.666667E-05 | global batch size: 32 | lm loss: 1.186456E+01 | load_balancing_loss: 8.779247E-01 | loss scale: 1.0 | grad norm: 3.868 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:22:47] iteration 3/ 10 | consumed samples: 96 | elapsed time per iteration (ms): 3115.0 | throughput per GPU (TFLOP/s/GPU): 123.0 | learning rate: 1.000000E-04 | global batch size: 32 | lm loss: 1.004934E+01 | load_balancing_loss: 8.783792E-01 | loss scale: 1.0 | grad norm: 3.355 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:22:50] iteration 4/ 10 | consumed samples: 128 | elapsed time per iteration (ms): 3122.0 | throughput per GPU (TFLOP/s/GPU): 122.8 | learning rate: 8.681980E-05 | global batch size: 32 | lm loss: 7.994447E+00 | load_balancing_loss: 8.822277E-01 | loss scale: 1.0 | grad norm: 3.496 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:22:53] iteration 5/ 10 | consumed samples: 160 | elapsed time per iteration (ms): 3131.2 | throughput per GPU (TFLOP/s/GPU): 122.4 | learning rate: 5.500000E-05 | global batch size: 32 | lm loss: 6.498296E+00 | load_balancing_loss: 8.777701E-01 | loss scale: 1.0 | grad norm: 2.971 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
saving checkpoint at iteration 5 to ./ckpts/1f1b_dsv2_lite_ep2_pp2_dp2 in torch_dist format
[2025-05-09 06:23:15] iteration 6/ 10 | consumed samples: 192 | elapsed time per iteration (ms): 3382.8 | throughput per GPU (TFLOP/s/GPU): 113.3 | learning rate: 2.318020E-05 | global batch size: 32 | lm loss: 5.168231E+00 | load_balancing_loss: 8.803471E-01 | loss scale: 1.0 | grad norm: 2.060 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:23:18] iteration 7/ 10 | consumed samples: 224 | elapsed time per iteration (ms): 3130.8 | throughput per GPU (TFLOP/s/GPU): 122.4 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 4.637459E+00 | load_balancing_loss: 8.775471E-01 | loss scale: 1.0 | grad norm: 1.641 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:23:22] iteration 8/ 10 | consumed samples: 256 | elapsed time per iteration (ms): 3139.1 | throughput per GPU (TFLOP/s/GPU): 122.1 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 4.372826E+00 | load_balancing_loss: 8.785877E-01 | loss scale: 1.0 | grad norm: 1.498 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:23:25] iteration 9/ 10 | consumed samples: 288 | elapsed time per iteration (ms): 3118.0 | throughput per GPU (TFLOP/s/GPU): 122.9 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 4.313786E+00 | load_balancing_loss: 8.782349E-01 | loss scale: 1.0 | grad norm: 1.477 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:23:28] iteration 10/ 10 | consumed samples: 320 | elapsed time per iteration (ms): 3113.2 | throughput per GPU (TFLOP/s/GPU): 123.1 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 4.122254E+00 | load_balancing_loss: 8.785514E-01 | loss scale: 1.0 | grad norm: 1.470 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
# mcore-0.11 resume 6-10 steps
[2025-05-09 06:25:45] iteration 6/ 10 | consumed samples: 192 | elapsed time per iteration (ms): 23436.3 | throughput per GPU (TFLOP/s/GPU): 16.4 | learning rate: 2.318020E-05 | global batch size: 32 | lm loss: 5.168231E+00 | load_balancing_loss: 8.803471E-01 | loss scale: 1.0 | grad norm: 2.060 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:25:49] iteration 7/ 10 | consumed samples: 224 | elapsed time per iteration (ms): 3629.7 | throughput per GPU (TFLOP/s/GPU): 105.6 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 4.637669E+00 | load_balancing_loss: 8.775379E-01 | loss scale: 1.0 | grad norm: 1.641 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:25:52] iteration 8/ 10 | consumed samples: 256 | elapsed time per iteration (ms): 3124.7 | throughput per GPU (TFLOP/s/GPU): 122.7 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 4.372930E+00 | load_balancing_loss: 8.786353E-01 | loss scale: 1.0 | grad norm: 1.499 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:25:55] iteration 9/ 10 | consumed samples: 288 | elapsed time per iteration (ms): 3117.4 | throughput per GPU (TFLOP/s/GPU): 122.9 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 4.312852E+00 | load_balancing_loss: 8.782726E-01 | loss scale: 1.0 | grad norm: 1.476 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
[2025-05-09 06:25:58] iteration 10/ 10 | consumed samples: 320 | elapsed time per iteration (ms): 3587.2 | throughput per GPU (TFLOP/s/GPU): 106.8 | learning rate: 1.000000E-05 | global batch size: 32 | lm loss: 4.126917E+00 | load_balancing_loss: 8.785437E-01 | loss scale: 1.0 | grad norm: 1.472 | num zeros: 0 | number of skipped iterations: 0 | number of nan iterations: 0 |
```
**Environment (please complete the following information):**
- Megatron-LM tag core-0.12 and core-0.11
- TransformerEngine-v1.13
- Image: nvcr.io/nvidia/pytorch:24.07-py3
**Proposed fix**
**Additional context**
Contributor guide
Assessment
This issue has not been assessed yet.