-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathprepare_glue_data.py
More file actions
172 lines (131 loc) · 5.62 KB
/
Copy pathprepare_glue_data.py
File metadata and controls
172 lines (131 loc) · 5.62 KB
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
#!/usr/bin/env python3
"""
准备 GLUE 数据集的脚本
用于下载和转换缺失的 GLUE 数据集(MRPC、STSB 等)
"""
import os
import sys
import pandas as pd
import urllib.request
import zipfile
import shutil
from pathlib import Path
def download_mrpc():
"""下载并准备 MRPC 数据集"""
print("正在准备 MRPC 数据集...")
data_dir = Path("/data/public_datasets/attention/glue_mrpc")
data_dir.mkdir(parents=True, exist_ok=True)
# MRPC 数据集的 URL
mrpc_train_url = "https://dl.fbaipublicfiles.com/senteval/senteval_data/msr_paraphrase_train.txt"
mrpc_test_url = "https://dl.fbaipublicfiles.com/senteval/senteval_data/msr_paraphrase_test.txt"
temp_dir = Path("/tmp/mrpc_download")
temp_dir.mkdir(parents=True, exist_ok=True)
try:
# 下载训练集
print(" 下载训练集...")
train_file = temp_dir / "msr_paraphrase_train.txt"
urllib.request.urlretrieve(mrpc_train_url, train_file)
# 下载测试集
print(" 下载测试集...")
test_file = temp_dir / "msr_paraphrase_test.txt"
urllib.request.urlretrieve(mrpc_test_url, test_file)
# 处理训练集
print(" 处理训练集...")
train_data = pd.read_csv(train_file, sep='\t', on_bad_lines='skip')
# 划分训练集和验证集(80/20 分割)
train_size = int(0.8 * len(train_data))
train_df = train_data[:train_size]
val_df = train_data[train_size:]
# 保存训练集
train_df.to_csv(data_dir / "train.tsv", sep='\t', index=False)
print(f" 训练集保存: {len(train_df)} 条样本")
# 保存验证集
val_df.to_csv(data_dir / "val.tsv", sep='\t', index=False)
print(f" 验证集保存: {len(val_df)} 条样本")
# 处理测试集
print(" 处理测试集...")
test_data = pd.read_csv(test_file, sep='\t', on_bad_lines='skip')
test_data.to_csv(data_dir / "test.tsv", sep='\t', index=False)
print(f" 测试集保存: {len(test_data)} 条样本")
print("✓ MRPC 数据集准备完成")
# 清理临时文件
shutil.rmtree(temp_dir)
except Exception as e:
print(f"✗ MRPC 数据集准备失败: {e}")
print("\n请手动下载 MRPC 数据集:")
print("1. 访问 https://www.microsoft.com/en-us/download/details.aspx?id=52398")
print("2. 下载数据集并放置在 /data/public_datasets/attention/glue_mrpc/")
print("3. 确保文件名为: train.tsv, val.tsv, test.tsv")
return False
return True
def download_stsb():
"""准备 STS-B 数据集"""
print("正在准备 STS-B 数据集...")
data_dir = Path("/data/public_datasets/attention/glue_stsb")
data_dir.mkdir(parents=True, exist_ok=True)
# STS-B 数据集需要从 GLUE 官方下载
print(" STS-B 数据集需要从 GLUE 官方网站下载")
print(" 请访问: https://gluebenchmark.com/tasks")
print(" 或使用 Hugging Face datasets 库:")
print(" from datasets import load_dataset")
print(" dataset = load_dataset('glue', 'stsb')")
return False
def check_existing_datasets():
"""检查现有数据集的状态"""
print("\n检查 GLUE 数据集状态:")
print("=" * 60)
datasets = {
'glue_cola': 'CoLA (单句分类)',
'glue_sst2': 'SST-2 (情感分析)',
'glue_mrpc': 'MRPC (句子对相似度)',
'glue_qnli': 'QNLI (问答 NLI)',
'glue_qqp': 'QQP (问题对相似度)',
'glue_mnli': 'MNLI (多类 NLI)',
'glue_rte': 'RTE (文本蕴含)',
'glue_stsb': 'STS-B (语义相似度)',
}
base_path = Path("/data/public_datasets/attention")
for dataset_name, description in datasets.items():
dataset_path = base_path / dataset_name
if not dataset_path.exists():
status = "✗ 目录不存在"
elif not any(dataset_path.iterdir()):
status = "✗ 目录为空"
else:
files = list(dataset_path.glob("*.tsv")) + list(dataset_path.glob("*.jsonl"))
has_train = any("train" in f.name for f in files)
has_val = any("val" in f.name for f in files)
has_test = any("test" in f.name for f in files)
if has_train and has_val and has_test:
status = "✓ 完整"
else:
missing = []
if not has_train: missing.append("train")
if not has_val: missing.append("val")
if not has_test: missing.append("test")
status = f"⚠ 缺少: {', '.join(missing)}"
print(f"{dataset_name:15} ({description:20}): {status}")
print("=" * 60)
def main():
print("GLUE 数据集准备工具")
print("=" * 60)
# 检查现有数据集
check_existing_datasets()
print("\n开始准备缺失的数据集...")
# 尝试下载 MRPC
mrpc_path = Path("/data/public_datasets/attention/glue_mrpc")
if not mrpc_path.exists() or not any(mrpc_path.iterdir()):
download_mrpc()
else:
print("✓ MRPC 数据集已存在")
# 尝试准备 STS-B
stsb_path = Path("/data/public_datasets/attention/glue_stsb")
if not stsb_path.exists() or not any(stsb_path.iterdir()):
download_stsb()
else:
print("✓ STS-B 数据集已存在")
print("\n" + "=" * 60)
print("准备完成!再次检查数据集状态:")
check_existing_datasets()
if __name__ == "__main__":
main()