#!/bin/bash

## 硬件相关环境设置
# CUDA
export CUDA_DEVICE_MAX_CONNECTIONS=32
export TORCH_NCCL_AVOID_RECORD_STREAMS=1
export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True
export NVTE_BWD_LAYERNORM_SM_MARGIN=16

## 软件相关环境变量(python / pytorch / megatron)
export OMP_NUM_THREADS=10
export PYTORCH_ENABLE_SAME_RAND_A100=1
export MHA_BWD_NO_ATOMIC_F64=1
export MAX_JOBS=20
export PROJECT_PATH=${PROJECT_PATH:-/root/megatron-infinigence}
export PYTHONPATH=${PROJECT_PATH}:${PROJECT_PATH}/Megatron-LM:$PYTHONPATH:/workspace/flash-attention
export ARCH="NVIDIA_H100"


## 软件相关环境变量(python / pytorch / megatron)
##############

# Distributed training variables
NNODES=$NNODES
GPUS_PER_NODE=8
GPU_NUM=$((${GPUS_PER_NODE}*${NNODES}))
WORLD_SIZE=$((${GPUS_PER_NODE}*${NNODES}))
NODE_RANK=$NODE_RANK
MASTER_ADDR=${MASTER_ADDR:-localhost}
MASTER_PORT=${MASTER_PORT:-7000}
SEED=${SEED:-1234}
GROUP_GEMM=${GROUP_GEMM:-0}
GROUP_GEMM_TE=${GROUP_GEMM_TE:-0}
TD=${TD:-allgather_fused}
RECOMPUTE=${RECOMPUTE:-0}
CPU_ADAM=${CPU_ADAM:-0}
SCHED=${SCHED:-1f1b}
USE_FLASH3=${USE_FLASH3:-0}


## ORIG:
NUM_LAYERS=${NUM_LAYERS:-61}
HIDDEN_SIZE=7168
NUM_ATTN_HEADS=128
FFN_HIDDEN_SIZE=18432
MOE_FFN_HIDDEN_SIZE=2048
MAX_POSITION_EMBEDDINGS=163840
EXTRA_VOCAB_SIZE=467
RMS_NORM_EPS=1e-6
MAX_SEQ_LEN=4096
MAX_PAD_LEN=4096
MICRO_BATCH_SIZE=1
GLOBAL_BATCH_SIZE=${GLOBAL_BATCH_SIZE:-4096}


# mla
QK_NOPE_HEAD_DIM=128
QK_ROPE_HEAD_DIM=64
V_HEAD_DIM=128
ROPE_THETA=10000
SCALE_FACTOR=40
ORIG_MAX_POSITION_EMBEDDINGS=4096
Q_LORA_RANK=1536
KV_LORA_RANK=512


# moe
NUM_EXPERTS=${NUM_EXPERTS:-256}
ROUTER_TOPK=8
NUM_SHARED_EXPERTS=1
FIRST_K_DENSE_REPLACE=${FIRST_K_DENSE_REPLACE:-3}


## trains
TRAIN_TOKENS=1000000000
WARMUP_TOKENS=10000
TRAIN_ITERS=$(( ${TRAIN_TOKENS} / ${GLOBAL_BATCH_SIZE} / ${MAX_SEQ_LEN} ))
LR_WARMUP_ITERS=$(( ${WARMUP_TOKENS}  / ${GLOBAL_BATCH_SIZE} / ${MAX_SEQ_LEN} ))
LR_DECAY_ITERS=$(( ${TRAIN_TOKENS} /  ${GLOBAL_BATCH_SIZE} / ${MAX_SEQ_LEN} ))
LR=5e-6
MIN_LR=1e-6

# Paths
SRC_PATH=$PROJECT_PATH/megatron_infini/pretrain_deepseek.py
DATA_PATH=${DATA_PATH:-/workspace/datasets/DeepSeek-V3/deepseekv3_text_document}
TOKENIZER_PATH=/workspace/DeepSeek-V3


PARALLEL_PERFORMANCE_ARGS=" \
    --bf16 \
    --use-distributed-optimizer \
    --sequence-parallel \
    --transformer-impl transformer_engine \
    --use-flash-attn \
    --tensor-model-parallel-size ${TP} \
    --tensor-model-parallel-for-expert ${TPE} \
    --expert-model-parallel-size ${EP} \
    --pipeline-model-parallel-size ${PP} \
    --moe-token-dispatcher-type-patch ${TD} \
    --mla-replicate-l1 \
    --no-bias-swiglu-fusion \
    "

if [ "$USE_FLASH3" -eq 1 ]; then
    PARALLEL_PERFORMANCE_ARGS="$PARALLEL_PERFORMANCE_ARGS \
        --module-impl {\"core_attn\":\"local\"} \
    "
fi

if [ "$GROUP_GEMM" -eq 1 ]; then
    PARALLEL_PERFORMANCE_ARGS="$PARALLEL_PERFORMANCE_ARGS \
        --moe-grouped-gemm \
    "
elif [ "$GROUP_GEMM_TE" -eq 1 ]; then
    PARALLEL_PERFORMANCE_ARGS="$PARALLEL_PERFORMANCE_ARGS \
        --moe-grouped-gemm-te \
    "
fi


if [ -n "$FORCE_DROP_AND_PADDING" ]; then
    PARALLEL_PERFORMANCE_ARGS="$PARALLEL_PERFORMANCE_ARGS \
        --force-drop-and-padding \
    "
fi


if [ -n "$VPP" ]; then
    PARALLEL_PERFORMANCE_ARGS="$PARALLEL_PERFORMANCE_ARGS \
        --num-layers-per-virtual-pipeline-stage ${VPP} \
    "
fi


if [ "$CUDA_DEVICE_MAX_CONNECTIONS" -gt 1 ]; then
    PARALLEL_PERFORMANCE_ARGS="$PARALLEL_PERFORMANCE_ARGS \
        --no-async-tensor-model-parallel-allreduce \
    "
fi


if [ "$SCHED" != "1f1b" ]; then
    PARALLEL_PERFORMANCE_ARGS="$PARALLEL_PERFORMANCE_ARGS \
        --scheduler ${SCHED} \
    "
fi


if [ -n "$HPP" ]; then
    if [ "$SCHED" == "1f1b" ]; then
        PARALLEL_PERFORMANCE_ARGS="$PARALLEL_PERFORMANCE_ARGS \
            --hetero-pipeline-stages $HPP \
        "
    else
        PARALLEL_PERFORMANCE_ARGS="$PARALLEL_PERFORMANCE_ARGS \
            --virtual-hetero-pipeline-stages $HPP \
        "
    fi
fi


if [ "$RECOMPUTE" == "1" ]; then
    PARALLEL_PERFORMANCE_ARGS="$PARALLEL_PERFORMANCE_ARGS \
        --recompute-method uniform \
        --recompute-num-layers 1 \
        --recompute-granularity full \
    "
elif [ "$MOE_LAYER_RECOMPUTE" == "1" ]; then
    PARALLEL_PERFORMANCE_ARGS="$PARALLEL_PERFORMANCE_ARGS \
        --moe-layer-recompute \
    "
elif [ "$MLP_RECOMPUTE" == "1" ]; then
    PARALLEL_PERFORMANCE_ARGS="$PARALLEL_PERFORMANCE_ARGS \
        --mlp-recompute \
    "
fi


if [ "$CPU_ADAM" == "1" ]; then
    PARALLEL_PERFORMANCE_ARGS="$PARALLEL_PERFORMANCE_ARGS \
        --optimizer hybridadam \
        --optimizer-offload-policy static \
        --optimizer-offload-fraction 1 \
        --optimizer-enable-pin \
    "
fi


if [ "$DP_OVERLAP" == "1" ]; then
    PARALLEL_PERFORMANCE_ARGS="$PARALLEL_PERFORMANCE_ARGS \
        --overlap-grad-reduce \
    "
fi


MOE_ARGS=" \
    --num-experts ${NUM_EXPERTS} \
    --moe-router-topk ${ROUTER_TOPK} \
    --moe-ffn-hidden-size ${MOE_FFN_HIDDEN_SIZE} \
    --enable-shared-expert \
    --num-shared-experts ${NUM_SHARED_EXPERTS} \
    --first-k-dense-replace ${FIRST_K_DENSE_REPLACE} \
    --moe-router-deepseekv3 \
    --moe-router-score-function sigmoid \
    --moe-router-enable-expert-bias True \
    --moe-router-topk-scaling-factor 2.5 \
    --moe-router-bias-update-rate 1e-3 \
    --moe-aux-loss-coeff 1e-4 \
    "

MLA_ARGS=" \
    --q-lora-rank ${Q_LORA_RANK} \
    --kv-lora-rank ${KV_LORA_RANK} \
    --qk-nope-head-dim ${QK_NOPE_HEAD_DIM} \
    --qk-rope-head-dim ${QK_ROPE_HEAD_DIM} \
    --v-head-dim ${V_HEAD_DIM} \
    --kv-channels ${V_HEAD_DIM} \
    --qk-layernorm \
    "


OTHER_NETWORK_ARGS=" \
    --use-mcore-models \
    --disable-bias-linear \
    --patch-tokenizer-type DeepSeekTokenizer \
    --tokenizer-model ${TOKENIZER_PATH} \
    --extra-vocab-size ${EXTRA_VOCAB_SIZE} \
    --max-padding-length ${MAX_PAD_LEN} \
    --swiglu \
    --normalization RMSNorm \
    --norm-epsilon ${RMS_NORM_EPS} \
    --use-rotary-position-embeddings \
    --no-rope-fusion \
    --position-embedding-type rope \
    --untie-embeddings-and-output-weights \
    --rotary-base ${ROPE_THETA} \
    --rotary-scaling-factor ${SCALE_FACTOR} \
    --rotary-seq-len-interpolation-factor 1 \
    --rotary-mscale 1.0 \
    --rotary-mscale-all-dim 1.0 \
    --original-max-position-embeddings ${ORIG_MAX_POSITION_EMBEDDINGS} \
    --rotary-beta-fast 32 \
    --rotary-beta-slow 1 \
    "


NETWORK_SIZE_ARGS="  \
    --num-layers ${NUM_LAYERS} \
    --hidden-size ${HIDDEN_SIZE} \
    --num-attention-heads ${NUM_ATTN_HEADS} \
    --ffn-hidden-size ${FFN_HIDDEN_SIZE} \
    --max-position-embeddings ${MAX_POSITION_EMBEDDINGS} \
    --seq-length ${MAX_SEQ_LEN} \
    "

TRAINING_ARGS=" \
    --micro-batch-size ${MICRO_BATCH_SIZE} \
    --global-batch-size ${GLOBAL_BATCH_SIZE} \
    --train-iters ${TRAIN_ITERS} \
    --eval-interval 10000 \
    --eval-iters 0 \
    "

LEARNING_ARGS=" \
    --lr ${LR} \
    --min-lr ${MIN_LR} \
    --lr-decay-style cosine \
    --lr-decay-iters ${LR_DECAY_ITERS} \
    --lr-warmup-iters ${LR_WARMUP_ITERS} \
    --attention-dropout 0.0 \
    --hidden-dropout 0.0 \
    --weight-decay 0.1 \
    --adam-beta1 0.9 \
    --adam-beta2 0.95 \
    --clip-grad 1.0 \
    --init-method-std 0.008 \
    --seed ${SEED} \
    "

LOAD_SAVE_ARGS=" \
    --no-load-optim \
    --no-load-rng \
    --num-workers 8 \
    --no-save-optim \
    "

if [ -n "$LOAD_PATH" ]; then
    LOAD_SAVE_ARGS="$LOAD_SAVE_ARGS \
        --load ${LOAD_PATH} \
    "
fi


if [ -n "$SAVE_PATH" ]; then
    LOAD_SAVE_ARGS="$LOAD_SAVE_ARGS \
        --save ${SAVE_PATH}
    "
fi


LOGGING_ARGS=" \
    --log-interval 1 \
    --log-throughput \
    --save-interval 500 \
    --timing-log-level 2 \
    "


INTERVAL_ARGS=" \
    --save-interval 10000 \
    --eval-interval 1000 \
    --eval-iters -1 \
    "


DATASET_ARGS=" \
    --num-workers 8 \
    --data-path ${DATA_PATH} \
    --split 99,1,0 \
    --dataset LLama-Pretrain-Idxmap \
    "


LAUNCHER=" \
    torchrun \
    --nproc_per_node ${GPUS_PER_NODE} \
    --nnodes ${NNODES} \
    --node_rank ${NODE_RANK} \
    --master_addr ${MASTER_ADDR} \
    --master_port ${MASTER_PORT} \
    "


RUN_CMD="${LAUNCHER} ${SRC_PATH} \
    ${PARALLEL_PERFORMANCE_ARGS} \
    ${MIXED_PRECISION_ARGS} \
    ${MOE_ARGS} \
    ${MLA_ARGS} \
    ${OTHER_NETWORK_ARGS} \
    ${NETWORK_SIZE_ARGS} \
    ${TRAINING_ARGS} \
    ${LEARNING_ARGS} \
    ${LOAD_SAVE_ARGS} \
    ${LOGGING_ARGS} \
    ${INTERVAL_ARGS} \
    ${DATASET_ARGS} \
    "

echo ${RUN_CMD}
