File size: 10,074 Bytes
932bc69
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
#!/usr/bin/env bash
# Infinity-Parser2.1-Flash 在线 DFlash2 训练。
# 平台两个节点各执行一次:每节点前 2 卡推理、后 6 卡训练,共 12 个 DDP rank。
set -Eeuo pipefail

# ============ 路径与平台分布式配置 ============
WS="${ROOT:-/inspire/sfs/project/inf-multimodal/public/wumengke}"
REPO="${REPO:-$WS/speculators}"
ENV_REPO="${ENV_REPO:-$WS/speculators}"
MODEL="${MODEL:-/inspire/sfs/project/inf-multimodal/public/data_mllm/publish_models/Infinity-Parser2.1-Flash-2608}"
DATA_ROOT="${DATA_ROOT:-$WS/datasets/infinity_parsers2_v2_1_max32768_vocab32k}"
DATA_DIR="${DATA_DIR:-$DATA_ROOT/dflash_data/full}"
# DFlash2 selector 使用完整的 248320-token 词表;数据目录已改为全词表映射。
OUTPUT_DIR="${OUTPUT_DIR:-${RUN_DIR:-$WS/model_weights/dflash2_parser2_1_flash_2node}}"
NNODES="${PET_NNODES:?需要 PET_NNODES=2}"
NODE_RANK="${PET_NODE_RANK:?需要 PET_NODE_RANK=0 或 1}"
[[ "$NNODES" == 2 && "${PET_NPROC_PER_NODE:-}" == 8 ]] || {
    echo "需要 2 节点、每节点 8 卡(2 推理 + 6 训练)" >&2; exit 1;
}
[[ "$NODE_RANK" == 0 || "$NODE_RANK" == 1 ]] || {
    echo "PET_NODE_RANK 必须为 0 或 1" >&2; exit 1;
}
DIST_MASTER_ADDR="${MASTER_ADDR:-${PET_MASTER_ADDR:?需要 MASTER_ADDR 或 PET_MASTER_ADDR}}"
DIST_MASTER_PORT="${MASTER_PORT:-${PET_MASTER_PORT:?需要 MASTER_PORT 或 PET_MASTER_PORT}}"
# 保留平台 NCCL/GLOO 网络设置,与现有两节点配方保持一致。
export NCCL_CROSS_NIC="${NCCL_CROSS_NIC:-0}"
if [[ -z "${GLOO_SOCKET_IFNAME:-}" && -n "${NCCL_SOCKET_IFNAME:-}" ]]; then
    export GLOO_SOCKET_IFNAME="$NCCL_SOCKET_IFNAME"
fi
RUN_NAME="${RUN_NAME:-dflash2-parser2_1-2node}"
SAVE_DIR="${CHECKPOINT_DIR:-$OUTPUT_DIR/$RUN_NAME/checkpoints}"
LOG_DIR="${LOG_DIR:-$OUTPUT_DIR/$RUN_NAME}"
VLLM_LOG="$LOG_DIR/vllm_node${NODE_RANK}.log"
TRAIN_LOG="$LOG_DIR/train_node${NODE_RANK}.log"

IFS=',' read -r -a GPU_LIST <<< "${CUDA_VISIBLE_DEVICES:-0,1,2,3,4,5,6,7}"
[[ ${#GPU_LIST[@]} == 8 ]] || { echo "每节点需要 8 张可见 GPU" >&2; exit 1; }
VLLM_GPUS=$(IFS=,; printf '%s' "${GPU_LIST[*]:0:2}")
TRAIN_GPUS=$(IFS=,; printf '%s' "${GPU_LIST[*]:2:6}")
NUM_TRAIN_GPUS=6

# DFlash2 的卷积、selector 和 CE/DPACE 配方沿用两节点参考脚本。
NUM_LAYERS="${NUM_LAYERS:-5}"
BLOCK_SIZE="${BLOCK_SIZE:-16}"
MAX_ANCHORS="${MAX_ANCHORS:-1024}"
DECAY_GAMMA=7
CONV_KERNEL_SIZE="${CONV_KERNEL_SIZE:-2}"
CONV_GROUP_SIZE="${CONV_GROUP_SIZE:-16}"
SELECTOR_RANK="${SELECTOR_RANK:-256}"
SELECTOR_TOP_K="${SELECTOR_TOP_K:-16}"
SELECTOR_LOSS_ALPHA="${SELECTOR_LOSS_ALPHA:-0.1}"
EPOCHS="${EPOCHS:-3}"
LR="${LR:-1e-4}"
MUON_LR=2e-4

# Parser2 使用 16K packing、2048 非因果滑窗及自己的辅助层/MRoPE 配置。
PACK_SEQ_LEN=16384
TARGET_LAYER_IDS=(2 7 12 17 22)
# 数据加载、HTTP 连接和请求重试使用 speculators 默认值。
# API 进程数由 launch_vllm.py 按可用 CPU 自动选择。
VLLM_MM_PROCESSOR_CACHE_GB="${VLLM_MM_PROCESSOR_CACHE_GB:-0}"
MEDIA_ROOT="/inspire/sfs/project/inf-multimodal/public"
VLLM_PORT="${VLLM_PORT:-8200}"
VLLM_ENDPOINT="http://127.0.0.1:${VLLM_PORT}/v1"

SPEC_PYTHON="${SPEC_PYTHON:-$ENV_REPO/speculators_venv/bin/python}"
TORCHRUN="${TORCHRUN:-$ENV_REPO/speculators_venv/bin/torchrun}"
VLLM_PYTHON="${VLLM_PYTHON:-$ENV_REPO/vllm_venv/bin/python}"
LAUNCH_VLLM="${LAUNCH_VLLM:-$REPO/scripts/launch_vllm.py}"
TRAIN_SCRIPT="${TRAIN_SCRIPT:-$REPO/scripts/train.py}"
export PYTHONPATH="$REPO/src:$REPO/hs_connectors/src${PYTHONPATH:+:$PYTHONPATH}"
export PYTHONUNBUFFERED=1
export NO_PROXY="${NO_PROXY:+$NO_PROXY,}127.0.0.1,localhost"
export no_proxy="${no_proxy:+$no_proxy,}127.0.0.1,localhost"
export HF_ENDPOINT="${HF_ENDPOINT:-https://hf-mirror.com}"
export HF_HOME="${HF_HOME:-$WS/.cache/huggingface}"
export HF_DATASETS_CACHE="${HF_DATASETS_CACHE:-$WS/datasets/.cache}"
export XDG_CACHE_HOME="${XDG_CACHE_HOME:-$WS/.cache}"
export TORCH_HOME="${TORCH_HOME:-$WS/.cache/torch}"
# 清掉 .bashrc 的共享 Triton 路径,由 PyTorch 自动按 GPU 选择 Triton 缓存。
# 节点之间只隔离缓存根目录;TRITON_HOME 保证直接调用 Triton 时也不写入 $HOME。
unset TRITON_CACHE_DIR
export TORCHINDUCTOR_CACHE_DIR="${TORCHINDUCTOR_CACHE_DIR:-$WS/.cache/torchinductor}/$RUN_NAME/node${NODE_RANK}"
export TRITON_HOME="$TORCHINDUCTOR_CACHE_DIR"
export VLLM_CACHE_ROOT="${VLLM_CACHE_ROOT:-$WS/.cache/vllm}"
export WANDB_PROJECT="${WANDB_PROJECT:-infinity-parser2-flash}"
export WANDB_MODE="${WANDB_MODE:-online}"
if [[ "$WANDB_MODE" == online && -z "${WANDB_API_KEY:-}" ]]; then
    WANDB_KEY_FILE="${WANDB_KEY_FILE:-$WS/.secrets/wandb_key}"
    [[ -s "$WANDB_KEY_FILE" ]] || { echo "缺少 W&B key:$WANDB_KEY_FILE" >&2; exit 1; }
    export WANDB_API_KEY="$(tr -d '[:space:]' < "$WANDB_KEY_FILE")"
fi

# ============ 准备数据与输出目录 ============
for path in "$MODEL/config.json" "$DATA_DIR/state.json" "$DATA_DIR/dataset_info.json"; do
    [[ -f "$path" ]] || { echo "缺少文件:$path" >&2; exit 1; }
done
mkdir -p "$SAVE_DIR" "$LOG_DIR"
exec 9>"$SAVE_DIR/training.lock.node${NODE_RANK}"
flock -n 9 || { echo "本节点已有任务使用 $SAVE_DIR" >&2; exit 1; }
HS_PATH="$(mktemp -d "/tmp/dflash2_parser2_node${NODE_RANK}.XXXXXX")"
VLLM_PID=""
TRAIN_PID=""

terminate_group() {
    local pid="$1"
    [[ -n "$pid" ]] || return 0
    kill -TERM -- "-$pid" 2>/dev/null || true
    for _ in {1..30}; do
        kill -0 -- "-$pid" 2>/dev/null || break
        sleep 1
    done
    kill -KILL -- "-$pid" 2>/dev/null || true
    wait "$pid" 2>/dev/null || true
}

cleanup() {
    local status=$?
    trap - EXIT INT TERM HUP
    terminate_group "$TRAIN_PID"
    terminate_group "$VLLM_PID"
    rm -r -- "$HS_PATH"  # 只删除本次 mktemp 创建的 hidden-state 目录。
    exit "$status"
}
trap cleanup EXIT
trap 'exit 130' INT
trap 'exit 143' TERM HUP

"$SPEC_PYTHON" - "$VLLM_PORT" <<'PY'
import socket
import sys

with socket.socket() as sock:
    sock.settimeout(1)
    if sock.connect_ex(("127.0.0.1", int(sys.argv[1]))) == 0:
        raise SystemExit(f"Port {sys.argv[1]} is already in use")
PY

# ============ 每节点独立的 vLLM 服务(TP=1 / DP=2) ============
echo "Node $NODE_RANK: vLLM GPUs=$VLLM_GPUS, training GPUs=$TRAIN_GPUS"
echo "Model: $MODEL"
echo "Data: $DATA_DIR"
echo "Draft vocab: full verifier vocabulary (248320 tokens)"
echo "Checkpoints: $SAVE_DIR"
echo "vLLM log: $VLLM_LOG"
echo "Training log: $TRAIN_LOG"
setsid env \
    -u RANK \
    -u WORLD_SIZE \
    -u LOCAL_RANK \
    -u LOCAL_WORLD_SIZE \
    -u MASTER_ADDR \
    -u MASTER_PORT \
    CUDA_VISIBLE_DEVICES="$VLLM_GPUS" \
    "$VLLM_PYTHON" "$LAUNCH_VLLM" "$MODEL" \
    --target-layer-ids "${TARGET_LAYER_IDS[@]}" \
    --hidden-states-backend file \
    --hidden-states-path "$HS_PATH" \
    -- \
    --tensor-parallel-size 1 \
    --data-parallel-size 2 \
    --data-parallel-backend mp \
    --nnodes 1 \
    --node-rank 0 \
    --master-addr 127.0.0.1 \
    --data-parallel-address 127.0.0.1 \
    --gpu-memory-utilization 0.9 \
    --max-model-len 65536 \
    --mm-processor-cache-gb "$VLLM_MM_PROCESSOR_CACHE_GB" \
    --served-model-name "$MODEL" \
    --allowed-local-media-path "$MEDIA_ROOT" \
    --limit-mm-per-prompt '{"image":16}' \
    --host 127.0.0.1 \
    --port "$VLLM_PORT" \
    >>"$VLLM_LOG" 2>&1 &
VLLM_PID=$!

echo "Waiting for local vLLM..."
deadline=$((SECONDS + 1800))
until curl \
    --noproxy '*' \
    -fsS \
    --connect-timeout 2 \
    --max-time 5 \
    "http://127.0.0.1:${VLLM_PORT}/health" >/dev/null 2>&1; do
    if ! kill -0 "$VLLM_PID" 2>/dev/null; then
        tail -n 100 "$VLLM_LOG" >&2
        echo "本机 vLLM 在就绪前退出" >&2
        exit 1
    fi
    if (( SECONDS >= deadline )); then
        echo "等待 vLLM 超过 1800 秒,见 $VLLM_LOG" >&2
        exit 1
    fi
    sleep 2
done

# ============ 两节点 DDP 训练(global world size = 12) ============
setsid env \
    -u RANK \
    -u WORLD_SIZE \
    -u LOCAL_RANK \
    -u LOCAL_WORLD_SIZE \
    CUDA_VISIBLE_DEVICES="$TRAIN_GPUS" \
    "$TORCHRUN" \
    --nnodes "$NNODES" \
    --node_rank "$NODE_RANK" \
    --nproc_per_node "$NUM_TRAIN_GPUS" \
    --master_addr "$DIST_MASTER_ADDR" \
    --master_port "$DIST_MASTER_PORT" \
    --rdzv_backend static \
    --rdzv_conf timeout=3600 \
    "$TRAIN_SCRIPT" \
    --verifier-name-or-path "$MODEL" \
    --data-path "$DATA_DIR" \
    --save-path "$SAVE_DIR" \
    --speculator-type dflash2 \
    --draft-arch qwen3 \
    --draft-hidden-act silu \
    --num-layers "$NUM_LAYERS" \
    --mask-token-id 248077 \
    --block-size "$BLOCK_SIZE" \
    --max-anchors "$MAX_ANCHORS" \
    --target-layer-ids "${TARGET_LAYER_IDS[@]}" \
    --draft-mrope-full-head-hack \
    --sliding-window 2048 \
    --sliding-window-non-causal \
    --draft-attn-impl simple_flex_attention \
    --loss-fn ce \
    --per-position-loss-weight dpace \
    --dflash-decay-gamma "$DECAY_GAMMA" \
    --conv-kernel-size "$CONV_KERNEL_SIZE" \
    --conv-group-size "$CONV_GROUP_SIZE" \
    --selector-rank "$SELECTOR_RANK" \
    --selector-top-k "$SELECTOR_TOP_K" \
    --selector-loss-alpha "$SELECTOR_LOSS_ALPHA" \
    --optimizer muon \
    --muon-lr "$MUON_LR" \
    --lr "$LR" \
    --scheduler-type cosine \
    --scheduler-warmup-ratio 0.01 \
    --epochs "$EPOCHS" \
    --checkpoint-freq 0.1 \
    --total-seq-len "$PACK_SEQ_LEN" \
    --train-data-ratio 0.99 \
    --noise-std 0 \
    --hidden-states-dtype bfloat16 \
    --hidden-states-backend file \
    --hidden-states-path "$HS_PATH" \
    --vllm-endpoint "$VLLM_ENDPOINT" \
    --on-missing generate \
    --on-generate delete \
    --seed 42 \
    --logger wandb \
    --log-dir "$LOG_DIR" \
    --run-name "$RUN_NAME" \
    >>"$TRAIN_LOG" 2>&1 &
TRAIN_PID=$!

status=0
wait -n -p finished_pid "$TRAIN_PID" "$VLLM_PID" || status=$?
if [[ "$finished_pid" == "$VLLM_PID" ]]; then
    tail -n 100 "$VLLM_LOG" >&2
    echo "训练期间本机 vLLM 退出" >&2
    exit 1
fi
if (( status != 0 )); then
    tail -n 100 "$TRAIN_LOG" >&2
    exit "$status"
fi
echo "Done. Checkpoints saved to $SAVE_DIR"