diff --git a/ant_pipeline/inference/run_serve_eagle3_qwen3.sh b/ant_pipeline/inference/run_serve_eagle3_qwen3.sh new file mode 100644 index 000000000..9f2d75442 --- /dev/null +++ b/ant_pipeline/inference/run_serve_eagle3_qwen3.sh @@ -0,0 +1,67 @@ +set -ex + +USE_EAGLE3=true # false or true +ENGINE_TYPE=trt # trt or sglang + +# for sglang +speculative_num_steps=5 +speculative_eagle_topk=8 +speculative_num_draft_tokens=4 + +# for trt +max_draft_len=4 + +# base_model=/mnt/modelops/models/Qwen3-32B +# eagle_model=/mnt/modelops/models/AngelSlim/Qwen3-32B_eagle3 + +# base_model=/mnt/modelops/models/Qwen3-14B +# eagle_model=/mnt/modelops/models/AngelSlim/Qwen3-14B_eagle3 + +base_model=/mnt/modelops/models/Qwen3-30B-A3B/ +eagle_model=/mnt/modelops/train/eagle3/output/qwen3-30B-A3b-eagle3_nnodes_8/epoch_9 + + +if [ ${ENGINE_TYPE} == "sglang" ];then + command=( + python3 -m sglang.launch_server \ + --model ${base_model} \ + --host 127.0.0.1 \ + --port 9122 \ + --mem-fraction 0.85 \ + --cuda-graph-max-bs 64 \ + --max-running-requests 64 \ + --chunked-prefill-size 8192 \ + --tp-size 8 + ) + + if $USE_EAGLE3; then + command+=( + --speculative-algorithm EAGLE3 \ + --speculative-draft-model-path ${eagle_model} \ + --speculative-num-steps ${speculative_num_steps} \ + --speculative-eagle-topk ${speculative_eagle_topk} \ + --speculative-num-draft-tokens ${speculative_num_draft_tokens} + ) + fi +else + command=( + python3 /mnt/modelops/460695/opensource/Ant-TensorRT-LLM/tensorrt_llm/commands/trtllm-serve.py \ + --model_path ${base_model} \ + --port 9122 \ + --host 127.0.0.1 \ + --backend pytorch \ + --max_batch_size 16 \ + --max_num_tokens 8192 \ + ) + + if $USE_EAGLE3; then + command+=(--spec_algo eagle3 --draft_model_path ${eagle_model} --max_draft_len ${max_draft_len}) + fi + +fi + +"${command[@]}" + + + + diff --git a/ant_pipeline/inference/run_sglang_benchmark.sh b/ant_pipeline/inference/run_sglang_benchmark.sh new file mode 100644 index 000000000..396d9aee2 --- /dev/null +++ b/ant_pipeline/inference/run_sglang_benchmark.sh @@ -0,0 +1,21 @@ +set -ex + +# pip install sglang[all] +target_model_path=/mnt/modelops/models/Qwen3-30B-A3B/ +draft_model_path=/mnt/modelops/train/eagle3/output/qwen3-30B-A3b-eagle3_nnodes_8/epoch_9 + +python3 -m sglang.launch_server \ + --model $target_model_path \ + --speculative-algorithm EAGLE3 \ + --speculative-draft-model-path $draft_model_path \ + --speculative-num-steps 3 \ + --speculative-eagle-topk 1 \ + --speculative-num-draft-tokens 4 \ + --mem-fraction-static 0.85 \ + --cuda-graph-max-bs 32 \ + --tp 2 \ + --context-length 8192 \ + --trust-remote-code \ + --host 127.0.0.1 \ + --port 9122 \ + --dtype bfloat16 \ No newline at end of file diff --git a/ant_pipeline/train/prepare_datasets.sh b/ant_pipeline/train/prepare_datasets.sh new file mode 100644 index 000000000..1f03cb135 --- /dev/null +++ b/ant_pipeline/train/prepare_datasets.sh @@ -0,0 +1,25 @@ +#!/bin/bash +set -ex + +script_dir=$(cd "$(dirname "$0")" && pwd) + +# 数据集列表 +datasets=( + "ultrachat:/mnt/modelops/datasets/HuggingFaceH4/ultrachat_200k/" +) + +# 循环处理每个数据集 +for entry in "${datasets[@]}"; do + dataset_name="${entry%%:*}" # 冒号前的部分 + data_path="${entry#*:}" # 冒号后的部分 + + output_path="/mnt/modelops/datasets/specforge_postprocess_${dataset_name}" + + rm -rf "${output_path}" + mkdir -p "${output_path}" + + python3 "${script_dir}/../../scripts/prepare_data.py" \ + --dataset "${dataset_name}" \ + --output_path "${output_path}" \ + --data-path "${data_path}" +done \ No newline at end of file diff --git a/ant_pipeline/train/prepare_env.sh b/ant_pipeline/train/prepare_env.sh new file mode 100644 index 000000000..8d2bb10f3 --- /dev/null +++ b/ant_pipeline/train/prepare_env.sh @@ -0,0 +1,22 @@ +set -ex + +script_dir=$(cd "$(dirname "$0")" && pwd) + +pip install --upgrade pip setuptools wheel +pip install --trusted-host pypi.org --trusted-host files.pythonhosted.org --use-pep517 ninja +pip install jinja2==3.1.2 + +pushd $script_dir/../../ +# pip install -r requirements.txt # 默认环境已装好 +pip install -e . +popd + +# 全量编译 安装 flash-attn +# MAX_JOBS=$(nproc) pip install flash-attn --no-build-isolation --no-cache-dir --force-reinstall # 默认环境已装好 + +# 全量编译 安装 deepspeed +# MAX_JOBS=$(nproc) \ +# TORCH_CUDA_ARCH_LIST="9.0" \ +# DS_BUILD_OP=1 \ +# TORCH_EXTENSIONS_DIR=~/.cache/torch_extensions \ +# pip install deepspeed --force-reinstall --global-option="build_ext" --global-option="-j192" diff --git a/ant_pipeline/train/train_by_deepspeed.sh b/ant_pipeline/train/train_by_deepspeed.sh new file mode 100644 index 000000000..ad0d6003e --- /dev/null +++ b/ant_pipeline/train/train_by_deepspeed.sh @@ -0,0 +1,83 @@ +#!/bin/bash +set -ex + +# 获取路径 +SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" &>/dev/null && pwd) +ROOT_DIR=$(dirname "$SCRIPT_DIR")/../ + +pushd ${ROOT_DIR} +pip install -e . +popd + +# 设置环境变量 +export TORCH_CUDA_ARCH_LIST="9.0" +export PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True +export TORCH_EXTENSIONS_DIR=~/.cache/torch_extensions +export MAX_JOBS=$(nproc) +export DS_BUILD_FUSED_ADAM=1 #只编译FUSED_ADAM 优化器 + +# 获取训练配置 +model_idx=${1} +nnodes=${2} + +if [ $nnodes -eq 1 ]; then + MASTER_ADDR=localhost + MASTER_PORT=$(( ( RANDOM % 10000 ) + 20000 )) + RANK=0 +else + export OMP_NUM_THREADS=8 + export NCCL_NET=IB + export NCCL_IB_DISABLE=0 + export NCCL_SOCKET_IFNAME=^docker0,lo,eth0 + export NCCL_IB_GID_INDEX=3 + export NCCL_DEBUG=0 +fi + +config_file=${SCRIPT_DIR}/train_configs.json +config=$(jq -r ".\"$model_idx\"" "$config_file") + +# 读取参数 +target_model_path=$(jq -r '.target_model_path' <<< "$config") +draft_model_config=$(jq -r '.draft_model_config' <<< "$config") +batch_size=$(jq -r '.batch_size' <<< "$config") +zero_stage=$(jq -r '.zero_stage' <<< "$config") + +# 设置相关路径参数 +train_name=$(basename "$draft_model_config" .json) +save_path=/mnt/modelops/train/eagle3/ +log_path=/mnt/modelops/train/eagle3/logs +output_path=/mnt/modelops/train/eagle3/output +mkdir -p ${log_path} +mkdir -p ${output_path} + +output_dir=${output_path}/${train_name}_nnodes_${nnodes} +train_log_path=${log_path}/${train_name}_nnodes_${nnodes}_rank${RANK}_train.log +nvidia_smi_log_path=${log_path}/${train_name}_nnodes_${nnodes}_rank${RANK}_nvidia_smi.log + +# 添加GPU 监控 +nvidia-smi --query-gpu=timestamp,index,utilization.gpu,memory.used,memory.total --format=csv -l 1 > ${nvidia_smi_log_path} & + +# 正式开启训练 +torchrun \ + --nnodes=$nnodes \ + --nproc_per_node=8 \ + --node_rank=$RANK \ + --master_addr=$MASTER_ADDR \ + --master_port=$MASTER_PORT \ + $ROOT_DIR/scripts/train_eagle3_online_deepspeed.py \ + --target-model-path ${target_model_path} \ + --draft-model-config ${ROOT_DIR}/configs/${draft_model_config} \ + --train-data-path /mnt/modelops/datasets/specforge_postprocess_ultrachat/ultrachat.jsonl \ + --output-dir $output_dir \ + --num-epochs 10 \ + --batch-size ${batch_size} \ + --learning-rate 1e-4 \ + --max-length 2048 \ + --chat-template qwen \ + --cache-dir ${save_path}/cache/ \ + --embedding-key model.embed_tokens.weight \ + --ttt-length 7 \ + --zero-stage ${zero_stage} \ + 2>&1 | tee $train_log_path + +echo "Training completed successfully. Output saved to $output_dir" \ No newline at end of file diff --git a/ant_pipeline/train/train_configs.json b/ant_pipeline/train/train_configs.json new file mode 100644 index 000000000..9040d07d7 --- /dev/null +++ b/ant_pipeline/train/train_configs.json @@ -0,0 +1,65 @@ +{ + "1": + { + "target_model_path": "/mnt/modelops/models/Qwen3-32B", + "draft_model_config": "qwen3-32b-eagle3.json", + "batch_size": 1, + "zero_stage": 1 + }, + "2": + { + "target_model_path": "/mnt/modelops/models/Qwen3-8B", + "draft_model_config": "qwen3-8b-eagle3.json", + "batch_size": 6, + "zero_stage": 1 + }, + "3": + { + "target_model_path": "/mnt/modelops/models/Qwen3-14B", + "draft_model_config": "qwen3-14b-eagle3.json", + "batch_size": 4, + "zero_stage": 1 + }, + "4": + { + "target_model_path": "/mnt/modelops/models/Qwen2.5-72B", + "draft_model_config": "qwen2.5-72b-eagle3.json", + "batch_size": 2, + "zero_stage": 3 + }, + "5": + { + "target_model_path": "/mnt/modelops/models/Qwen3-30B-A3B", + "draft_model_config": "qwen3-30B-A3b-eagle3.json", + "batch_size": 1, + "zero_stage": 1 + }, + "6": + { + "target_model_path": "/mnt/modelops/models/Qwen3-235B-A22B-Instruct-2507", + "draft_model_config": "qwen3-235B-A22B-eagle3.json", + "batch_size": 1, + "zero_stage": 3 + }, + "7": + { + "target_model_path": "/mnt/modelops/487922/deepseek-ai__DeepSeek-V2-Lite-Chat", + "draft_model_config": "deepseek-v2-lite-chat-eagle3.json", + "batch_size": 4, + "zero_stage": 1 + }, + "8": + { + "target_model_path": "/mnt/modelops/models/QwQ-32B", + "draft_model_config": "qwq-32b-eagle3.json", + "batch_size": 2, + "zero_stage": 1 + }, + "9": + { + "target_model_path": "/mnt/modelops/models/Qwen3-4B/", + "draft_model_config": "qwen3-4b-eagle3.json", + "batch_size": 12, + "zero_stage": 1 + } +} \ No newline at end of file diff --git a/configs/deepseek-v2-lite-chat-eagle3.json b/configs/deepseek-v2-lite-chat-eagle3.json new file mode 100644 index 000000000..60adda506 --- /dev/null +++ b/configs/deepseek-v2-lite-chat-eagle3.json @@ -0,0 +1,39 @@ +{ + "architectures": [ + "LlamaForCausalLMEagle3" + ], + "attention_bias": false, + "attention_dropout": 0.0, + "bos_token_id": 100000, + "eos_token_id": 100001, + "head_dim": 128, + "hidden_act": "silu", + "hidden_size": 2048, + "initializer_range": 0.02, + "intermediate_size": 10944, + "max_position_embeddings": 163840, + "max_window_layers": 64, + "model_type": "llama", + "num_attention_heads": 16, + "num_hidden_layers": 1, + "num_key_value_heads": 16, + "rms_norm_eps": 1e-06, + "rope_scaling": { + "beta_fast": 32.0, + "beta_slow": 1.0, + "factor": 40.0, + "mscale": 0.707, + "mscale_all_dim": 0.707, + "original_max_position_embeddings": 4096, + "type": "yarn" + }, + "rope_theta": 10000, + "sliding_window": null, + "tie_word_embeddings": false, + "torch_dtype": "bfloat16", + "transformers_version": "4.33.1", + "use_cache": true, + "use_sliding_window": false, + "vocab_size": 102400, + "draft_vocab_size": 32000 +} \ No newline at end of file diff --git a/configs/qwen2.5-72b-eagle3.json b/configs/qwen2.5-72b-eagle3.json new file mode 100644 index 000000000..2fe165bd9 --- /dev/null +++ b/configs/qwen2.5-72b-eagle3.json @@ -0,0 +1,28 @@ +{ + "architectures": [ + "LlamaForCausalLMEagle3" + ], + "attention_dropout": 0.0, + "bos_token_id": 151643, + "eos_token_id": 151645, + "hidden_act": "silu", + "hidden_size": 8192, + "initializer_range": 0.02, + "intermediate_size": 29568, + "max_position_embeddings": 32768, + "max_window_layers": 70, + "model_type": "llama", + "num_attention_heads": 64, + "num_hidden_layers": 1, + "num_key_value_heads": 8, + "rms_norm_eps": 1e-06, + "rope_theta": 1000000.0, + "sliding_window": 131072, + "tie_word_embeddings": false, + "torch_dtype": "bfloat16", + "transformers_version": "4.43.1", + "use_cache": true, + "use_sliding_window": false, + "vocab_size": 152064, + "draft_vocab_size": 32000 +} diff --git a/configs/qwen3-30B-A3B-eagle3.json b/configs/qwen3-14b-eagle3.json similarity index 61% rename from configs/qwen3-30B-A3B-eagle3.json rename to configs/qwen3-14b-eagle3.json index 558cb1804..048d502fe 100644 --- a/configs/qwen3-30B-A3B-eagle3.json +++ b/configs/qwen3-14b-eagle3.json @@ -1,31 +1,31 @@ { "architectures": [ "LlamaForCausalLMEagle3" - ], + ], + "draft_vocab_size": 32000, + "num_hidden_layers": 1, + "model_type": "llama", "attention_bias": false, "attention_dropout": 0.0, "bos_token_id": 151643, "eos_token_id": 151645, "head_dim": 128, "hidden_act": "silu", - "hidden_size": 2048, + "hidden_size": 5120, "initializer_range": 0.02, - "intermediate_size": 12288, - "max_position_embeddings": 2048, - "max_window_layers": 48, - "model_type": "llama", - "num_attention_heads": 32, - "num_hidden_layers": 1, - "num_key_value_heads":4, + "intermediate_size": 17408, + "max_position_embeddings": 40960, + "max_window_layers": 40, + "num_attention_heads": 40, + "num_key_value_heads": 8, "rms_norm_eps": 1e-06, "rope_scaling": null, - "rope_theta": 1000000.0, + "rope_theta": 1000000, "sliding_window": null, "tie_word_embeddings": false, "torch_dtype": "bfloat16", - "transformers_version": "4.53.2", + "transformers_version": "4.51.0", "use_cache": true, "use_sliding_window": false, - "vocab_size": 151936, - "draft_vocab_size": 32000 -} + "vocab_size": 151936 +} \ No newline at end of file diff --git a/configs/qwen3-30B-A3b-eagle3.json b/configs/qwen3-30B-A3b-eagle3.json new file mode 100644 index 000000000..d8c68ecfa --- /dev/null +++ b/configs/qwen3-30B-A3b-eagle3.json @@ -0,0 +1,39 @@ +{ + "architectures": [ + "LlamaForCausalLMEagle3" + ], + "draft_vocab_size": 32000, + "num_hidden_layers": 1, + "model_type": "llama", + "attention_bias": false, + "attention_dropout": 0.0, + "bos_token_id": 151643, + "decoder_sparse_step": 1, + "eos_token_id": 151645, + "head_dim": 128, + "hidden_act": "silu", + "hidden_size": 2048, + "initializer_range": 0.02, + "intermediate_size": 6144, + "max_position_embeddings": 40960, + "max_window_layers": 48, + "mlp_only_layers": [], + "moe_intermediate_size": 768, + "norm_topk_prob": true, + "num_attention_heads": 32, + "num_experts": 128, + "num_experts_per_tok": 8, + "num_key_value_heads": 4, + "output_router_logits": false, + "rms_norm_eps": 1e-06, + "rope_scaling": null, + "rope_theta": 1000000.0, + "router_aux_loss_coef": 0.001, + "sliding_window": null, + "tie_word_embeddings": false, + "torch_dtype": "bfloat16", + "transformers_version": "4.51.0", + "use_cache": true, + "use_sliding_window": false, + "vocab_size": 151936 +} \ No newline at end of file diff --git a/configs/qwen3-4b-eagle3.json b/configs/qwen3-4b-eagle3.json index 41ae128fd..3cf30e9dc 100644 --- a/configs/qwen3-4b-eagle3.json +++ b/configs/qwen3-4b-eagle3.json @@ -1,7 +1,10 @@ { "architectures": [ "LlamaForCausalLMEagle3" - ], + ], + "draft_vocab_size": 32000, + "num_hidden_layers": 1, + "model_type": "llama", "attention_bias": false, "attention_dropout": 0.0, "bos_token_id": 151643, @@ -13,19 +16,16 @@ "intermediate_size": 9728, "max_position_embeddings": 40960, "max_window_layers": 36, - "model_type": "llama", "num_attention_heads": 32, - "num_hidden_layers": 1, "num_key_value_heads": 8, "rms_norm_eps": 1e-06, "rope_scaling": null, "rope_theta": 1000000, "sliding_window": null, - "tie_word_embeddings": false, + "tie_word_embeddings": true, "torch_dtype": "bfloat16", "transformers_version": "4.51.0", "use_cache": true, "use_sliding_window": false, - "vocab_size": 151936, - "draft_vocab_size": 32000 -} + "vocab_size": 151936 +} \ No newline at end of file diff --git a/configs/qwen3-8b-eagle3.json b/configs/qwen3-8b-eagle3.json index 2bf9844cb..f53c32c32 100644 --- a/configs/qwen3-8b-eagle3.json +++ b/configs/qwen3-8b-eagle3.json @@ -1,7 +1,10 @@ { "architectures": [ "LlamaForCausalLMEagle3" - ], + ], + "draft_vocab_size": 32000, + "num_hidden_layers": 1, + "model_type": "llama", "attention_bias": false, "attention_dropout": 0.0, "bos_token_id": 151643, @@ -11,12 +14,10 @@ "hidden_size": 4096, "initializer_range": 0.02, "intermediate_size": 12288, - "max_position_embeddings": 2048, + "max_position_embeddings": 40960, "max_window_layers": 36, - "model_type": "llama", "num_attention_heads": 32, - "num_hidden_layers": 1, - "num_key_value_heads":8 , + "num_key_value_heads": 8, "rms_norm_eps": 1e-06, "rope_scaling": null, "rope_theta": 1000000, @@ -26,6 +27,5 @@ "transformers_version": "4.51.0", "use_cache": true, "use_sliding_window": false, - "vocab_size": 151936, - "draft_vocab_size": 32000 -} + "vocab_size": 151936 +} \ No newline at end of file diff --git a/configs/qwq-32B-eagle3.json b/configs/qwq-32b-eagle3.json similarity index 96% rename from configs/qwq-32B-eagle3.json rename to configs/qwq-32b-eagle3.json index 8f7d7908d..11cd49299 100644 --- a/configs/qwq-32B-eagle3.json +++ b/configs/qwq-32b-eagle3.json @@ -11,7 +11,7 @@ "intermediate_size": 27648, "max_position_embeddings": 40960, "max_window_layers": 64, - "model_type": "qwen2", + "model_type": "llama", "num_attention_heads": 40, "num_hidden_layers": 1, "num_key_value_heads": 8, diff --git a/scripts/prepare_data.py b/scripts/prepare_data.py index 63baa7261..079d4c16f 100644 --- a/scripts/prepare_data.py +++ b/scripts/prepare_data.py @@ -126,7 +126,8 @@ def main(): args = parse_args() # load dataset if args.dataset == "ultrachat": - ds = load_dataset("HuggingFaceH4/ultrachat_200k")["train_sft"] + dataset_path = args.data_path if args.data_path is not None else "HuggingFaceH4/ultrachat_200k/" + ds = load_dataset(dataset_path)["train_sft"] proc_fn = process_ultrachat_row elif args.dataset == "sharegpt": if args.data_path is None: diff --git a/scripts/train_eagle3_online_deepspeed.py b/scripts/train_eagle3_online_deepspeed.py new file mode 100644 index 000000000..8b1e31e73 --- /dev/null +++ b/scripts/train_eagle3_online_deepspeed.py @@ -0,0 +1,444 @@ +import argparse +import hashlib +import os +from tqdm import tqdm + +import torch +import torch.distributed as dist +from torch.utils.data.distributed import DistributedSampler +from torch.utils.data import DataLoader + +from accelerate.utils import set_seed +from datasets import load_dataset + +from transformers import AutoModelForCausalLM, AutoTokenizer +from transformers.integrations import HfDeepSpeedConfig + +from specforge import AutoDraftModelConfig, AutoEagle3DraftModel, OnlineEagle3Model +from specforge.data import build_eagle3_dataset, generate_vocab_mapping_file +from specforge.data.utils import DataCollatorWithPadding +from specforge.utils import get_last_checkpoint, print_with_rank, rank_0_priority + +import deepspeed + + +def save_draft_model_safetensors( + model_engine, + output_dir: str, + exclude_substr=("embed",), + shard_size="2GB", +): + if dist.get_rank() == 0: + os.makedirs(output_dir, exist_ok=True) + + dist.barrier() + + draft = model_engine.module.draft_model + params_to_gather = [ + p for n, p in draft.named_parameters() + if not any(s in n.lower() for s in exclude_substr) + ] + + with deepspeed.zero.GatheredParameters(params_to_gather, modifier_rank=0), torch.no_grad(): + if dist.get_rank() == 0: + full_sd = { + k: v.detach().cpu() + for k, v in draft.state_dict().items() + if not any(s in k.lower() for s in exclude_substr) + } + draft.save_pretrained( + output_dir, + state_dict=full_sd, + safe_serialization=True, # -> safetensors + max_shard_size=shard_size + ) + print(f"[rank0] Draft model (filtered) saved to: {output_dir}") + + dist.barrier() + + +def get_zero_config(args,total_steps,warmup_steps): + zero_stages = { + "0":{ + "zero_optimization": { + "stage": 0, + "allgather_partitions": True, + "allgather_bucket_size": 5e8, + "overlap_comm": False, + "reduce_scatter": True, + "reduce_bucket_size": 5e8, + "contiguous_gradients": True, + "round_robin_gradients": True + } + }, + "1":{ + "zero_optimization": { + "stage": 1, + "overlap_comm": True, + "allgather_partitions": True, + "reduce_scatter": True, + "contiguous_gradients": True + } + }, + "2":{ + "zero_optimization": { + "stage": 2, + "overlap_comm": True, + "contiguous_gradients": True, + "reduce_scatter": True, + "reduce_bucket_size": 5e8, + "round_robin_gradients":True, + "allgather_partitions": True, + } + }, + "3":{ + "zero_optimization": { + "stage": 3, + "overlap_comm": True, + "allgather_partitions": True, + "reduce_scatter": True, + "contiguous_gradients": True, + "stage3_gather_16bit_weights_on_model_save":True, + "stage3_max_live_parameters" : 1e9, + "stage3_param_persistence_threshold": 1e6, + # "zero_hpz_partition_size": 8, + "stage3_max_reuse_distance" : 1e9, + "reduce_bucket_size": 5e8, + "stage3_prefetch_bucket_size" : 5e8, + } + } + } + ds_config = { + "train_micro_batch_size_per_gpu": args.batch_size, + "gradient_accumulation_steps": args.gradient_accumulation_steps, + "bf16": { + "enabled": True, + "immediate_grad_update": True, + "check_grad_overflow": True + }, + "fp16": { + "enabled": False + }, + "optimizer": { + "type": "AdamW", + "params": { + "lr": args.learning_rate, + "weight_decay": 0.01, + "betas": [0.9, 0.95] + } + }, + "scheduler": { + "type": "WarmupDecayLR", + "params": { + "total_num_steps": total_steps, + "warmup_num_steps": warmup_steps, + "warmup_min_lr":0, + "warmup_max_lr": args.learning_rate + } + }, + **zero_stages[str(args.zero_stage)], + "communication_data_type": "bf16", + } + + return ds_config + +def parse_args(): + parser = argparse.ArgumentParser(description="Train Eagle3 with online data") + parser.add_argument("--target-model-path", type=str, required=True) + parser.add_argument("--draft-model-config", type=str, required=True) + parser.add_argument( + "--embedding-key", + type=str, + default="model.embed_tokens.weight", + help="The key of the embedding weight to load from the target model", + ) + parser.add_argument("--train-data-path", type=str, required=True) + parser.add_argument("--eval-data-path", type=str, default=None) + parser.add_argument("--num-epochs", type=int, default=10) + parser.add_argument("--batch-size", type=int, default=1) + parser.add_argument("--learning-rate", type=float, default=1e-4) + parser.add_argument("--max-length", type=int, default=2048) + parser.add_argument("--warmup-ratio", type=float, default=0.02) + parser.add_argument( + "--ttt-length", + type=int, + default=7, + help="The length for Test-Time Training (TTT).", + ) + parser.add_argument("--chat-template", type=str, default="llama3") + parser.add_argument("--cache-key", type=str, default=None) + parser.add_argument("--cache-dir", type=str, default="./cache") + parser.add_argument("--output-dir", type=str, required=True) + parser.add_argument("--eval-interval", type=int, default=1) + parser.add_argument("--save-interval", type=int, default=1) + parser.add_argument("--seed", type=int, default=0) + parser.add_argument( + "--dist-timeout", + type=int, + default=20, + help="Timeout for collective communication in minutes", + ) + parser.add_argument("--attention-backend", type=str, default="flex_attention") + parser.add_argument("--resume", action="store_true") + parser.add_argument("--local_rank", type=int, default=0) + parser.add_argument("--gradient-accumulation-steps", type=int, default=1) + parser.add_argument("--zero-stage", type=int, default=3) + args = parser.parse_args() + return args + + +def print_on_rank0(message): + if dist.get_rank() == 0: + print(message) + + +def main(): + args = parse_args() + set_seed(args.seed) + + deepspeed.init_distributed() + + local_rank = int(os.getenv("LOCAL_RANK",0)) + rank = dist.get_rank() + world_size = dist.get_world_size() + local_world_size = int(os.getenv("LOCAL_WORLD_SIZE", 8)) + torch.cuda.set_device(local_rank) + print_with_rank(f"WORLD SIZE={world_size}-RANK={rank}-LOCAL_WORLD_SIZE={local_world_size}-LOCAL RANK={local_rank}") + print("Initialized distributed environment") + + draft_model_last_checkpoint = None + if args.resume and os.path.isdir(args.output_dir): + print_on_rank0(args.output_dir) + draft_model_last_checkpoint = get_last_checkpoint(args.output_dir) + print_on_rank0(f"Last checkpoint detected: {draft_model_last_checkpoint}") + + draft_model_config = AutoDraftModelConfig.from_file(args.draft_model_config) + if draft_model_last_checkpoint: + draft_model = ( + AutoEagle3DraftModel.from_pretrained( + draft_model_last_checkpoint, attention_backend=args.attention_backend + ) + .cuda() + .to(torch.bfloat16) + ) + else: + draft_model = ( + AutoEagle3DraftModel.from_config( + draft_model_config, attention_backend=args.attention_backend + ) + .cuda() + .to(torch.bfloat16) + ) + + draft_model.load_embedding(args.target_model_path, embedding_key=args.embedding_key) + draft_model.freeze_embedding() + print_with_rank("Initialized draft model") + + tokenizer = AutoTokenizer.from_pretrained(args.target_model_path) + cache_params_string = ( + f"{args.train_data_path}-" + f"{args.max_length}-" + f"{args.chat_template}-" + f"{args.target_model_path}" # Tokenizer may also different + ) + cache_key = hashlib.md5(cache_params_string.encode()).hexdigest() + train_dataset = load_dataset("json", data_files=args.train_data_path)["train"] + with rank_0_priority(): + train_eagle3_dataset = build_eagle3_dataset( + dataset=train_dataset, + tokenizer=tokenizer, + chat_template=args.chat_template, + max_length=args.max_length, + cache_dir=os.path.join(args.cache_dir, "processed_dataset"), + cache_key=cache_key, + ) + vocab_mapping_path = generate_vocab_mapping_file( + dataset=train_eagle3_dataset, + target_vocab_size=draft_model_config.vocab_size, + draft_vocab_size=draft_model_config.draft_vocab_size, + cache_dir=os.path.join(args.cache_dir, "vocab_mapping"), + cache_key=cache_key, + ) + draft_model.load_vocab_mapping(vocab_mapping_path) + print_with_rank("Loaded vocab mapping") + + train_sampler = DistributedSampler( + train_eagle3_dataset, + num_replicas=dist.get_world_size(), + rank=dist.get_rank(), + shuffle=True, + seed=args.seed, + drop_last=False, + ) + + train_dataloader = DataLoader( + train_eagle3_dataset, + batch_size=args.batch_size, + sampler=train_sampler, + num_workers=8, + pin_memory=True, + drop_last=False, + collate_fn=DataCollatorWithPadding(), + persistent_workers=True, + ) + + total_steps = args.num_epochs * len(train_dataloader) + warmup_steps = int(total_steps * args.warmup_ratio) + print_with_rank("Initialized train dataloader") + + ds_config = get_zero_config(args,total_steps,warmup_steps) + dschf = HfDeepSpeedConfig(ds_config) # keep this object alive + target_model = AutoModelForCausalLM.from_pretrained(args.target_model_path, torch_dtype=torch.bfloat16) + + # https://github.com/deepspeedai/DeepSpeed/issues/7461 + for m in target_model.modules(): + if "SparseMoeBlock" in m.__class__.__name__: + deepspeed.utils.set_z3_leaf_modules(target_model, [m.__class__]) + print(f"--------Setting zero3 leaf for model on class with name: {m.__class__.__name__}---------") + break + + for _, param in target_model.named_parameters(): + param.requires_grad = False + print_with_rank("Initialized target model") + + eagle3_model = OnlineEagle3Model( + target_model= target_model, + draft_model = draft_model, + length = args.ttt_length, + attention_backend = args.attention_backend, + ) + print_with_rank("Initialized Eagle3 model") + + model_engine, optimizer, _, scheduler = deepspeed.initialize( + model=eagle3_model, + config=ds_config, + model_parameters=eagle3_model.draft_model.parameters() + ) + print_with_rank("Initialized Eagle3 DeepSpeed model") + + print_with_rank("Start Training!") + start_epoch = 0 + if draft_model_last_checkpoint is not None: + print_on_rank0( + f"Resuming draft model training from checkpoint: {draft_model_last_checkpoint}" + ) + state_path = os.path.join(draft_model_last_checkpoint, "training_state.pt") + + if os.path.exists(state_path): + state = torch.load(state_path, map_location="cpu", weights_only=False) + + try: + model_engine.optimizer.load_state_dict(state["optimizer_state_dict"]) + print_on_rank0("Successfully loaded optimizer state_dict.") + except: + print_on_rank0("Warning: Failed to load optimizer state_dict.") + + try: + scheduler.load_state_dict(state["scheduler_state_dict"]) + print_on_rank0("Successfully loaded scheduler state_dict.") + except: + print_on_rank0("Warning: Failed to load scheduler state_dict.") + + start_epoch = state["epoch"] + 1 + print_on_rank0(f"Resuming from epoch {start_epoch}") + else: + print_on_rank0( + f"Warning: Checkpoint directory {draft_model_last_checkpoint} found, but training_state.pt is missing. Starting from scratch." + ) + + dist.barrier() + print_on_rank0(f"Starting training from epoch {start_epoch}") + + global_step = 0 + for epoch in range(start_epoch, args.num_epochs): + train_dataloader.sampler.set_epoch(epoch) + model_engine.train() + draft_model.train() # for consistency + + epoch_acces = [[] for _ in range(model_engine.module.length)] + epoch_plosses = [[] for _ in range(model_engine.module.length)] + + # Training loop + pbar = tqdm(train_dataloader, desc=f"Epoch {epoch}", disable=(local_rank != 0)) + for data in pbar: + input_ids = data["input_ids"].to(model_engine.device) + attention_mask = data["attention_mask"].to(model_engine.device) + loss_mask = data["loss_mask"].to(model_engine.device) + + # Forward + plosses, _, acces = model_engine( + input_ids=input_ids, + attention_mask=attention_mask, + loss_mask=loss_mask, + ) + # Weighted loss + ploss_weight = [0.8 ** i for i in range(len(plosses))] + ploss = sum(ploss_weight[i] * plosses[i] for i in range(len(plosses))) + + # Backward + model_engine.backward(ploss) + + model_engine.step() # zero_grad, optimizer.step, scheduler.step, grad clip + + global_step += 1 + + # Accumulate for epoch logging + epoch_acces = [epoch_acces[i] + [acces[i]] for i in range(len(acces))] + epoch_plosses = [ + epoch_plosses[i] + [plosses[i].item()] for i in range(len(plosses)) + ] + + # Update progress bar + pbar.set_postfix({"ploss": ploss.item(), "lr": model_engine.get_lr()[0]}) + + # Epoch-level logging + for i in range(len(epoch_acces)): + acc_i = torch.tensor(epoch_acces[i]).cuda().mean() + dist.all_reduce(acc_i) + acc_i = acc_i / dist.get_world_size() + acc_i = acc_i.item() + print_on_rank0( + f"Train Epoch [{epoch + 1}/{args.num_epochs}], position {i}, Acc: {acc_i:.4f}" + ) + + for i in range(len(epoch_plosses)): + loss_i = torch.tensor(epoch_plosses[i]).cuda().mean() + dist.all_reduce(loss_i) + loss_i = loss_i / dist.get_world_size() + loss_i = loss_i.item() + print_on_rank0( + f"Train Epoch [{epoch + 1}/{args.num_epochs}], position {i}, pLoss: {loss_i:.4f}" + ) + + # Save checkpoint + if epoch % args.save_interval == 0 or epoch == args.num_epochs - 1: + epoch_output_dir = os.path.join(args.output_dir, f"epoch_{epoch}") + if dist.get_rank() == 0: + os.makedirs(epoch_output_dir, exist_ok=True) + dist.barrier() + # Only save full state on rank 0 + if dist.get_rank() == 0: + # Save training state + state_to_save = { + "optimizer_state_dict": model_engine.optimizer.state_dict(), + "scheduler_state_dict": scheduler.state_dict(), + "epoch": epoch, + "args": args, + } + torch.save( + state_to_save, + os.path.join(epoch_output_dir, "training_state.pt"), + ) + print_on_rank0(f"Saved training state to {epoch_output_dir}/training_state.pt") + + save_draft_model_safetensors(model_engine, epoch_output_dir) + + dist.barrier() + + print_on_rank0("Training completed.") + if dist.is_initialized(): + dist.destroy_process_group() + + +if __name__ == "__main__": + main() \ No newline at end of file diff --git a/specforge/data/template.py b/specforge/data/template.py index 16b0f5d71..02d980997 100644 --- a/specforge/data/template.py +++ b/specforge/data/template.py @@ -113,3 +113,13 @@ def get_all_template_names(self) -> List[str]: end_of_turn_token="<|im_end|>\n", ), ) + +TEMPLATE_REGISTRY.register( + name="deepseek", + template=ChatTemplate( + assistant_header="Assistant:", + user_header="User:", + system_prompt="You are a helpful assistant.", + end_of_turn_token="", + ), +)