File size: 8,654 Bytes
bf314e8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""UMA demo YAML 配置解析器 - 把 YAML 转换为训练启动命令与相关变量。

用法:
    python _parse_config.py <config.yaml> <action> [hydra_config_path]

action:
    name           打印实验名称
    command        打印训练启动命令 (python/torchrun/srun 风格)
                   需要额外传入 hydra_config_path,作为 train.py -c 的目标
    env            打印环境变量 export 语句
    env-args       打印 env_setup.sh 的参数 (conda_env module1 module2 ...)
    data-files     打印 preflight_check.sh 要检查的路径
                   (checkpoint_location + data.train_dataset.splits.*.src + data.val_dataset.splits.*.src)
    slurm          打印 SLURM 配置变量
    hydra-config   把 YAML 剥离 demo meta 后的部分(纯 hydra 内容)打印到 stdout
                   run.sh 会把它重定向到 outputs/xxx_ts/hydra_config.yaml

顶层 meta 键(只给 demo 工具用,不进 hydra):
    name, description, launch, env, slurm, nccl

其余顶层键(data, job, runner, reducer, train_dataset, ...)作为 hydra 内容。
"""

import os
import sys

import yaml


# 这些顶层键只给 demo 工具用,在写 hydra_config 时会被剥离
META_KEYS = {"name", "description", "launch", "env", "slurm", "nccl"}


def load_config(path):
    with open(path) as f:
        return yaml.safe_load(f)


def get_train_py_path(config_path):
    """根据 demo/configs/xxx.yaml 的位置推算 UMA/train.py 路径。"""
    # config_path: .../UMA/demo/configs/xxx.yaml
    demo_dir = os.path.dirname(os.path.dirname(os.path.abspath(config_path)))
    uma_dir = os.path.dirname(demo_dir)
    return os.path.join(uma_dir, "train.py")


# -----------------------------------------------------------------------------
# action 实现
# -----------------------------------------------------------------------------

def print_name(cfg):
    print(cfg.get("name", "uma_train"))


def print_command(cfg, config_path, hydra_config_path):
    """拼 python/torchrun/srun train.py -c <hydra_config_path>。"""
    if not hydra_config_path:
        print(
            "[ERROR] command 动作需要提供 hydra_config_path 作为第三个参数",
            file=sys.stderr,
        )
        sys.exit(1)

    launch = cfg.get("launch", {}) or {}
    num_nodes = launch.get("num_nodes", 1)
    num_gpus = launch.get("num_gpus", 1)
    launcher = launch.get("launcher", "python")

    train_py = get_train_py_path(config_path)

    if num_nodes > 1:
        # 多节点: sbatch 内会用 srun 包裹 torchrun, 命令本体是 torchrun + train.py
        # 注意: master_addr/master_port/node_rank 在 run.sh 生成的 submit.sh 里
        # 基于 $SLURM_* 变量导出, 这里占位用 ${MASTER_ADDR} 等 shell 变量
        print(
            "torchrun \\\n"
            f"  --nnodes={num_nodes} \\\n"
            "  --node_rank=${SLURM_NODEID} \\\n"
            f"  --nproc_per_node={num_gpus} \\\n"
            "  --rdzv_id=${SLURM_JOB_ID} \\\n"
            "  --rdzv_backend=c10d \\\n"
            "  --rdzv_endpoint=${MASTER_ADDR}:${MASTER_PORT} \\\n"
            f"  {train_py} \\\n"
            f"  -c {hydra_config_path}"
        )
    elif launcher == "torchrun" or num_gpus > 1:
        # 单节点多卡: torchrun
        print(
            "torchrun \\\n"
            "  --nnodes=1 \\\n"
            f"  --nproc_per_node={num_gpus} \\\n"
            f"  {train_py} \\\n"
            f"  -c {hydra_config_path}"
        )
    else:
        # 单卡: 直接 python
        print(f"python {train_py} -c {hydra_config_path}")


def print_env(cfg):
    """打印环境变量 export 语句。"""
    launch = cfg.get("launch", {}) or {}
    num_nodes = launch.get("num_nodes", 1)
    num_gpus = launch.get("num_gpus", 1)

    omp = launch.get("omp_num_threads", 1)
    print(f"export OMP_NUM_THREADS={omp}")

    if num_gpus > 1:
        devices = ",".join(str(i) for i in range(num_gpus))
        print(f"export HIP_VISIBLE_DEVICES={devices}")
    elif num_gpus == 1:
        print("export HIP_VISIBLE_DEVICES=0")

    # 多节点 NCCL 配置
    if num_nodes > 1:
        nccl = cfg.get("nccl", {}) or {}
        print("export HSA_FORCE_FINE_GRAIN_PCIE=1")
        if nccl.get("socket_ifname"):
            print(f"export NCCL_SOCKET_IFNAME={nccl['socket_ifname']}")
        if nccl.get("ib_hca"):
            print(f"export NCCL_IB_HCA={nccl['ib_hca']}")
        if nccl.get("proto"):
            print(f"export NCCL_PROTO={nccl['proto']}")


def print_env_args(cfg):
    """打印 env_setup.sh 的参数: conda_env module1 module2 ..."""
    env = cfg.get("env", {}) or {}
    conda_env = env.get("conda_env", "chem")
    modules = env.get("modules", []) or []
    parts = [conda_env] + list(modules)
    print(" ".join(parts))


def _dig_src_paths(node, out):
    """递归找出 splits 下所有 src 字段。"""
    if isinstance(node, dict):
        if "splits" in node and isinstance(node["splits"], dict):
            for split_name, split_val in node["splits"].items():
                if isinstance(split_val, dict) and "src" in split_val:
                    out.append(split_val["src"])
        for v in node.values():
            _dig_src_paths(v, out)
    elif isinstance(node, list):
        for item in node:
            _dig_src_paths(item, out)


def _dig_checkpoint_locations(node, out):
    """递归找出 checkpoint_location 字段。"""
    if isinstance(node, dict):
        if "checkpoint_location" in node and isinstance(
            node["checkpoint_location"], str
        ):
            out.append(node["checkpoint_location"])
        for v in node.values():
            _dig_checkpoint_locations(v, out)
    elif isinstance(node, list):
        for item in node:
            _dig_checkpoint_locations(item, out)


def print_data_files(cfg):
    """打印 preflight 要检查的路径: checkpoint + 数据 src。"""
    paths = []
    data = cfg.get("data", {})
    _dig_src_paths(data, paths)
    _dig_checkpoint_locations(cfg, paths)
    # 去重保持顺序
    seen = set()
    for p in paths:
        if p and p not in seen and not str(p).startswith("??"):
            seen.add(p)
            print(p)


def print_slurm(cfg):
    """打印 SLURM 配置变量。"""
    launch = cfg.get("launch", {}) or {}
    slurm = cfg.get("slurm", {}) or {}
    num_nodes = launch.get("num_nodes", 1)
    num_gpus = launch.get("num_gpus", 1)

    name = cfg.get("name", "uma_train")
    partition = slurm.get("partition", "k100ai")
    time_limit = slurm.get("time", "12:00:00")
    cpus = slurm.get("cpus_per_task", 128)

    # UMA 的多节点策略: 每个节点 1 个 srun task, 由 torchrun 在该节点内再 spawn
    # nproc_per_node=num_gpus 个进程. 因此无论单/多节点 NTASKS_PER_NODE 都是 1.
    ntasks = 1
    cpus_per_task = cpus
    _ = num_gpus  # 保留 num_gpus 以便 GPUS_PER_NODE 正确填写

    print(f"JOB_NAME={name}")
    print(f"PARTITION={partition}")
    print(f"NODES={num_nodes}")
    print(f"NTASKS_PER_NODE={ntasks}")
    print(f"CPUS_PER_TASK={cpus_per_task}")
    print(f"GPUS_PER_NODE={num_gpus}")
    print(f"TIME={time_limit}")


def print_hydra_config(cfg):
    """把 cfg 剥离 META_KEYS 后的部分用 YAML 输出。"""
    hydra_only = {k: v for k, v in cfg.items() if k not in META_KEYS}
    # sort_keys=False 保留原顺序; default_flow_style=False 强制 block 风格
    yaml.safe_dump(
        hydra_only,
        sys.stdout,
        sort_keys=False,
        default_flow_style=False,
        allow_unicode=True,
    )


# -----------------------------------------------------------------------------
# main
# -----------------------------------------------------------------------------

def main():
    if len(sys.argv) < 3:
        print(__doc__, file=sys.stderr)
        sys.exit(1)

    config_path = sys.argv[1]
    action = sys.argv[2]
    hydra_config_path = sys.argv[3] if len(sys.argv) >= 4 else None

    cfg = load_config(config_path)

    actions = {
        "name": lambda: print_name(cfg),
        "command": lambda: print_command(cfg, config_path, hydra_config_path),
        "env": lambda: print_env(cfg),
        "env-args": lambda: print_env_args(cfg),
        "data-files": lambda: print_data_files(cfg),
        "slurm": lambda: print_slurm(cfg),
        "hydra-config": lambda: print_hydra_config(cfg),
    }

    fn = actions.get(action)
    if fn is None:
        print(f"未知 action: {action}", file=sys.stderr)
        print(f"可用: {', '.join(actions.keys())}", file=sys.stderr)
        sys.exit(1)
    fn()


if __name__ == "__main__":
    main()