File size: 6,743 Bytes
c34ff1f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Parser2 DFlash2 双节点入口:验证实际命令和完整词表,不启动 GPU 进程。"""

import argparse
import json
from pathlib import Path

import numpy as np
from torch.distributed.run import get_args_parser

from speculators.train.cli import parse_vocab_mappings
from speculators.train.config import TrainConfig
from tests.unit.scripts.test_qwen3_6_two_node_launch import (
    NETWORK_ENV,
    SCRIPT_DIR,
    assert_node_caches,
    assert_teacher_stopped,
    flag,
    read_capture,
    run_nodes,
)
from tests.unit.scripts.test_qwen3_6_two_node_launch import (
    launch_env as launch_env,  # noqa: PLC0414 复用模拟启动 fixture。
)

SCRIPT = SCRIPT_DIR / "dflash2_infinity_parser2_flash_online_2node.sh"


def test_parser2_two_node_launch(launch_env, monkeypatch):
    env = launch_env
    # 使用真实目录结构,确认传入的是 prepared Arrow 数据而不是 regen 根目录。
    env.pop("DATA_DIR")
    data = (
        Path(env["ROOT"])
        / "datasets/infinity_parsers2_v2_1_max32768_vocab32k/dflash_data/full"
    )
    data.mkdir(parents=True)
    for name in ("state.json", "dataset_info.json"):
        (data / name).write_text("{}")
    d2t = np.zeros(248320, dtype=np.int64)
    t2d = np.ones(248320, dtype=np.bool_)
    np.save(data / "d2t.npy", d2t)
    np.save(data / "t2d.npy", t2d)
    (Path(env["MODEL"]) / "config.json").write_text(
        json.dumps({"model_type": "qwen3_5", "text_config": {"vocab_size": 248320}})
    )
    env["NNODE_TEST_EXPECT_NODES"] = "2"
    for output, returncode in run_nodes("dflash2", env, (0, 1), script=SCRIPT):
        assert returncode == 0, output

    configs = []
    for rank in (0, 1):
        teacher = read_capture(env, "vllm", rank)
        train = read_capture(env, "train", rank)
        assert teacher["env"]["CUDA_VISIBLE_DEVICES"] == "7,6"
        assert train["env"]["CUDA_VISIBLE_DEVICES"] == "5,4,3,2,1,0"
        assert "VLLM_MEDIA_LOADING_THREAD_COUNT" not in teacher["env"]
        assert "--api-server-count" not in teacher["argv"]
        for record in (teacher, train):
            for name in ("RANK", "WORLD_SIZE", "LOCAL_RANK", "LOCAL_WORLD_SIZE"):
                assert name not in record["env"]
            for name, value in NETWORK_ENV.items():
                assert record["env"][name] == value
        assert "MASTER_ADDR" not in teacher["env"]
        assert "MASTER_PORT" not in teacher["env"]
        for name, value in {
            "--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",
            "--max-model-len": "65536",
            "--mm-processor-cache-gb": "0",
            "--served-model-name": env["MODEL"],
        }.items():
            assert flag(teacher["argv"], name) == value

        monkeypatch.setenv("PET_NPROC_PER_NODE", "8")
        distributed = get_args_parser().parse_args(train["argv"])
        assert distributed.nnodes == "2"
        assert distributed.nproc_per_node == "6"
        assert distributed.node_rank == rank
        assert distributed.master_addr == env["MASTER_ADDR"]
        assert str(distributed.master_port) == env["MASTER_PORT"]
        assert distributed.rdzv_backend == "static"
        assert not distributed.standalone
        assert not distributed.no_python

        cfg = TrainConfig.resolve(distributed.training_script_args).flatten()
        configs.append(cfg)
        assert cfg["run_name"] == "dflash2-parser2_1-2node"
        run_dir = (
            Path(env["ROOT"])
            / "model_weights/dflash2_parser2_1_flash_2node"
            / cfg["run_name"]
        )
        assert cfg["save_path"] == str(run_dir / "checkpoints")
        assert cfg["log_dir"] == str(run_dir)
        assert_node_caches(teacher, train, env, rank, cfg["run_name"])
        for name, value in {
            "speculator_type": "dflash2",
            "checkpoint_freq": 0.1,
            "verifier_name_or_path": env["MODEL"],
            "data_path": str(data),
            "draft_vocab_size": None,
            "draft_arch": "qwen3",
            "num_layers": 5,
            "mask_token_id": 248077,
            "target_layer_ids": [2, 7, 12, 17, 22],
            "draft_mrope_full_head_hack": True,
            "sliding_window": 2048,
            "sliding_window_non_causal": True,
            "full_attention_indices": [],
            "total_seq_len": 16384,
            "block_size": 16,
            "max_anchors": 1024,
            "sample_from_anchor": None,
            "loss_fn": "ce",
            "per_position_loss_weight": "dpace",
            "conv_kernel_size": 2,
            "conv_group_size": 16,
            "selector_rank": 256,
            "selector_top_k": 16,
            "selector_loss_alpha": 0.1,
            "num_workers": 12,
            "prefetch_factor": 4,
            "fetch_threads": 1,
            "dataloader_in_order": True,
            "vllm_http_keepalive": True,
            "request_timeout": 120,
            "max_retries": 3,
            "generation_validation_retries": 2,
            "max_consecutive_generation_failures": 20,
            "fail_on_hidden_state_error": False,
        }.items():
            assert cfg[name] == value, name
        assert cfg["vllm_endpoint"] == f"http://127.0.0.1:{env['VLLM_PORT']}/v1"
        loaded_d2t, loaded_t2d, vocab_size = parse_vocab_mappings(
            argparse.Namespace(**cfg)
        )
        assert vocab_size == 248320
        np.testing.assert_array_equal(loaded_d2t.numpy(), d2t)
        np.testing.assert_array_equal(loaded_t2d.numpy(), t2d)
        hs_path = Path(cfg["hidden_states_path"])
        assert hs_path.parent == Path("/tmp")
        assert not hs_path.exists()  # 脚本退出时清理本次目录。
        assert str(hs_path) == flag(teacher["argv"], "--hidden-states-path")
        assert_teacher_stopped(teacher)

    for key in ("save_path", "run_name", "log_dir"):
        assert configs[0][key] == configs[1][key]
    assert configs[0]["hidden_states_path"] != configs[1]["hidden_states_path"]
    # 训练直接加载全词表映射,不改写数据目录。
    np.testing.assert_array_equal(np.load(data / "d2t.npy"), d2t)
    np.testing.assert_array_equal(np.load(data / "t2d.npy"), t2d)


def test_missing_parser2_prepared_data_fails_before_teacher(launch_env):
    launch_env.pop("DATA_DIR")
    [(output, returncode)] = run_nodes("dflash2", launch_env, (0,), script=SCRIPT)
    assert returncode != 0
    assert "dflash_data/full/state.json" in output
    assert not list(Path(launch_env["NNODE_TEST_CAPTURE"]).iterdir())