-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathextract_results.py
More file actions
157 lines (129 loc) · 5.45 KB
/
Copy pathextract_results.py
File metadata and controls
157 lines (129 loc) · 5.45 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
#!/usr/bin/env python3
import os
import re
import math
from pathlib import Path
def extract_accuracy_from_log(log_file, is_glue=False):
"""从日志文件中提取准确率"""
if not os.path.exists(log_file):
return None
with open(log_file, 'r', encoding='utf-8') as f:
content = f.read()
if is_glue:
# GLUE 任务:提取验证集准确率(max val accuracy)
# 需要匹配 "max val accuracy: XX.XX%" 但排除后面的 "test accuracy: 0.00%"
pattern = r'max val accuracy: ([\d.]+)%; the test accuracy'
matches = re.findall(pattern, content)
if matches:
# 取最后一个(最终的 max val accuracy)
return float(matches[-1])
else:
# 文本分类任务:提取测试集准确率
pattern = r'the test accuracy corresponding to the max val accuracy: ([\d.]+)%'
matches = re.findall(pattern, content)
if matches:
# 取最后一个
return float(matches[-1])
return None
def get_results_for_dataset(dataset_name, base_dir="/home/lzr/code/new_attention/output"):
"""获取某个数据集的所有结果"""
is_glue = dataset_name.startswith('glue_')
# 确定方法名称
if is_glue:
base_method = f"transformer_{dataset_name}"
conv_method = f"conv_transformer_{dataset_name}"
else:
base_method = "base_transformer"
conv_method = "conv_base_transformer"
seeds = [10, 20, 30, 40, 50]
base_results = []
conv_results = []
for seed in seeds:
if is_glue:
base_log = os.path.join(base_dir, dataset_name, base_method, f"seed_{seed}", "log_rank0.txt")
conv_log = os.path.join(base_dir, dataset_name, conv_method, f"seed_{seed}_new_gate", "log_rank0.txt")
else:
base_log = os.path.join(base_dir, dataset_name, base_method, f"seed_{seed}", "log_rank0.txt")
conv_log = os.path.join(base_dir, dataset_name, conv_method, f"seed_{seed}_new_gate", "log_rank0.txt")
base_acc = extract_accuracy_from_log(base_log, is_glue)
conv_acc = extract_accuracy_from_log(conv_log, is_glue)
if base_acc is not None:
base_results.append(base_acc)
if conv_acc is not None:
conv_results.append(conv_acc)
return base_results, conv_results
def calculate_stats(values):
"""计算平均值和标准差"""
if not values:
return None, None
n = len(values)
mean = sum(values) / n
# 样本标准差
variance = sum((x - mean) ** 2 for x in values) / (n - 1) if n > 1 else 0.0
std = math.sqrt(variance)
return mean, std
def format_result(mean, std):
"""格式化结果为 mean ± std 形式"""
if mean is None or std is None:
return "N/A"
return f"{mean:.2f} ± {std:.2f}"
# 文本分类数据集
text_datasets = ['Rotten_Tomatoes', 'imdb', 'ag_news', '20_newsgroups']
# GLUE 数据集
glue_datasets = ['glue_cola', 'glue_wnli', 'glue_mnli', 'glue_mrpc', 'glue_qnli', 'glue_qqp', 'glue_rte', 'glue_sst2']
print("=" * 80)
print("文本分类数据集(测试集准确率)")
print("=" * 80)
results_text = {}
for dataset in text_datasets:
base_results, conv_results = get_results_for_dataset(dataset)
base_mean, base_std = calculate_stats(base_results)
conv_mean, conv_std = calculate_stats(conv_results)
results_text[dataset] = {
'base': (base_mean, base_std, base_results),
'conv': (conv_mean, conv_std, conv_results)
}
print(f"\n{dataset}:")
print(f" base_transformer: {format_result(base_mean, base_std)}")
print(f" conv_base_transformer: {format_result(conv_mean, conv_std)}")
if base_results:
print(f" base values: {base_results}")
if conv_results:
print(f" conv values: {conv_results}")
print("\n" + "=" * 80)
print("GLUE 数据集(验证集准确率)")
print("=" * 80)
results_glue = {}
for dataset in glue_datasets:
base_results, conv_results = get_results_for_dataset(dataset)
base_mean, base_std = calculate_stats(base_results)
conv_mean, conv_std = calculate_stats(conv_results)
results_glue[dataset] = {
'base': (base_mean, base_std, base_results),
'conv': (conv_mean, conv_std, conv_results)
}
print(f"\n{dataset}:")
print(f" transformer: {format_result(base_mean, base_std)}")
print(f" conv_transformer: {format_result(conv_mean, conv_std)}")
if base_results:
print(f" base values: {base_results}")
if conv_results:
print(f" conv values: {conv_results}")
# 生成 Markdown 格式的输出
print("\n" + "=" * 80)
print("Markdown 格式输出")
print("=" * 80)
print("\n### 文本分类数据集(测试集准确率)")
for dataset in text_datasets:
base_mean, base_std, _ = results_text[dataset]['base']
conv_mean, conv_std, _ = results_text[dataset]['conv']
print(f"\n### {dataset}")
print(f"- **base_transformer**: {format_result(base_mean, base_std)}%")
print(f"- **conv_base_transformer**: {format_result(conv_mean, conv_std)}%")
print("\n### GLUE 数据集(验证集准确率)")
for dataset in glue_datasets:
base_mean, base_std, _ = results_glue[dataset]['base']
conv_mean, conv_std, _ = results_glue[dataset]['conv']
print(f"\n### {dataset}")
print(f"- **transformer**: {format_result(base_mean, base_std)}%")
print(f"- **conv_transformer**: {format_result(conv_mean, conv_std)}%")