File size: 7,113 Bytes
83112d8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
01_data 数据流水线总控
按顺序执行所有数据处理阶段

阶段概览:
  1  噪声过滤        → data/1_noise_filtered/
  2  OOV过滤         → data/2_oov_filtered/
  3  精确去重         → data/3_dedup/
  4  格式预处理       → data/4_preprocessed/
  5  规则标注         → data/5_rule_labels/
  6  Qwen语义标注     → data/6_qwen_labels/     (需GPU)
  7  评测句提取       → data/7_eval_sentences/
  8A 采样-纯质量基线  → data/8_sample_A/train.txt
  8B 采样-配额+质量   → data/8_sample_B/train.txt
  8C 采样-加权采样    → data/8_sample_C/train.txt
  8D 采样-Embedding   → data/8_sample_D/train.txt  (需GPU)

用法:
    python scripts/01_data/run_all.py              # 跑所有阶段 (1-8A)
    python scripts/01_data/run_all.py --stage 5    # 只跑某一阶段
    python scripts/01_data/run_all.py --from 7     # 从某阶段开始跑
    python scripts/01_data/run_all.py --stage 8    # 跑所有 stage 8 采样 (A/B/C/D)
    python scripts/01_data/run_all.py --stage 8B   # 只跑 stage 8B
    python scripts/01_data/run_all.py --stage 8D   # 只跑 stage 8D (Embedding)
"""

import argparse
import subprocess
import sys
from pathlib import Path

ROOT = Path(__file__).parent.parent.parent

# ── 阶段定义 ────────────────────────────────────────────────────────────────────
# type="io"   : 标准 input/output 参数
# type="plain" : 无参数直接运行

STAGES = [
    {
        "id": "1", "name": "噪声过滤 (CHAT标注 / 时间戳 / 极短行)",
        "script": "stage1_noise_filter.py", "type": "io",
        "input": "data/raw_files", "output": "data/1_noise_filtered",
    },
    {
        "id": "2", "name": "OOV过滤 (低频噪声词过多的句子)",
        "script": "stage2_oov_filter.py", "type": "io",
        "input": "data/1_noise_filtered", "output": "data/2_oov_filtered",
    },
    {
        "id": "3", "name": "精确去重(全局 MD5)",
        "script": "stage3_dedup.py", "type": "io",
        "input": "data/2_oov_filtered", "output": "data/3_dedup",
    },
    {
        "id": "4", "name": "格式预处理(去说话人前缀 / 全大写规范 / 标题行过滤)",
        "script": "stage4_preprocess.py", "type": "io",
        "input": "data/3_dedup", "output": "data/4_preprocessed",
    },
    {
        "id": "5", "name": "规则标注(BLiMP 67子任务 / BLiMP Supplement / EWoK 11域)",
        "script": "stage5_rule_label.py", "type": "io",
        "input": "data/4_preprocessed", "output": "data/5_rule_labels",
    },
    {
        "id": "6", "name": "Qwen语义标注(9维度深度分析,需GPU)",
        "script": "stage6_qwen_label.py", "type": "io",
        "input": "data/4_preprocessed", "output": "data/6_qwen_labels",
    },
    {
        "id": "7", "name": "评测句提取",
        "script": "stage7a_extract_eval.py", "type": "plain",
    },
    {
        "id": "8A", "name": "采样方案A:纯质量基线",
        "script": "stage8_sample_A.py", "type": "plain",
    },
    {
        "id": "8B", "name": "采样方案B:任务配额 + 质量填充",
        "script": "stage8_sample_B.py", "type": "plain",
    },
    {
        "id": "8C", "name": "采样方案C:加权采样",
        "script": "stage8_sample_C.py", "type": "plain",
    },
    {
        "id": "8D", "name": "采样方案D:Embedding 相似度采样(需GPU)",
        "script": "stage8_sample_D.py", "type": "plain",
    },
]


def run_stage(stage):
    script = Path(__file__).parent / stage["script"]

    print(f"\n{'='*60}")
    print(f"阶段 {stage['id']}: {stage['name']}")
    print(f"  脚本: {script.name}")
    if stage["type"] == "io":
        print(f"  输入: {ROOT / stage['input']}")
        print(f"  输出: {ROOT / stage['output']}")
    print(f"{'='*60}")

    if stage["type"] == "io":
        cmd = [sys.executable, str(script),
               "--input_dir",  str(ROOT / stage["input"]),
               "--output_dir", str(ROOT / stage["output"])]
    else:
        cmd = [sys.executable, str(script)]

    result = subprocess.run(cmd, check=False)
    if result.returncode != 0:
        print(f"\n[ERROR] 阶段 {stage['id']} 失败,退出码 {result.returncode}")
        sys.exit(result.returncode)
    print(f"\n[OK] 阶段 {stage['id']} 完成")


def match_stages(selector: str) -> list[dict]:
    """根据选择器匹配阶段。支持: "5", "8", "8B", "8D" 等。"""
    selector = selector.upper()
    matched = []
    for s in STAGES:
        sid = s["id"].upper()
        if sid == selector:
            matched.append(s)
        elif selector.isdigit() and sid.startswith(selector):
            # "8" 匹配 "8A", "8B", "8C", "8D"
            matched.append(s)
    return matched


def main():
    parser = argparse.ArgumentParser(
        description="01_data 数据流水线总控",
        formatter_class=argparse.RawDescriptionHelpFormatter,
    )
    parser.add_argument("--stage", type=str, default=None,
                        help="只运行指定阶段(如 --stage 5, --stage 8, --stage 8B)")
    parser.add_argument("--from", type=str, default="1", dest="from_stage",
                        help="从指定阶段开始运行(如 --from 7)")
    parser.add_argument("--to", type=str, default=None, dest="to_stage",
                        help="运行到指定阶段为止(如 --to 6)")
    parser.add_argument("--list", action="store_true",
                        help="列出所有阶段")
    args = parser.parse_args()

    if args.list:
        print("可用阶段:")
        for s in STAGES:
            print(f"  {s['id']:4s}  {s['name']}")
        return

    if args.stage is not None:
        stages_to_run = match_stages(args.stage)
    else:
        # --from / --to 范围选择(只匹配纯数字前缀)
        from_idx = 0
        to_idx = len(STAGES)
        for i, s in enumerate(STAGES):
            num = s["id"].rstrip("ABCD")
            if num == args.from_stage:
                from_idx = i
                break
        if args.to_stage:
            for i, s in enumerate(STAGES):
                num = s["id"].rstrip("ABCD")
                if num == args.to_stage:
                    # 包含该阶段的所有子阶段
                    to_idx = i + 1
                    while to_idx < len(STAGES) and STAGES[to_idx]["id"].startswith(args.to_stage):
                        to_idx += 1
                    break
        stages_to_run = STAGES[from_idx:to_idx]

    if not stages_to_run:
        print(f"没有找到匹配的阶段: {args.stage or args.from_stage}")
        print("使用 --list 查看所有可用阶段")
        sys.exit(1)

    print(f"将运行 {len(stages_to_run)} 个阶段: {[s['id'] for s in stages_to_run]}")

    for stage in stages_to_run:
        run_stage(stage)

    print(f"\n{'='*60}")
    print("数据流水线完成!")
    print(f"{'='*60}")


if __name__ == "__main__":
    main()