-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathbalance_data.py
More file actions
212 lines (177 loc) · 9.94 KB
/
Copy pathbalance_data.py
File metadata and controls
212 lines (177 loc) · 9.94 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
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
import os
import json
import logging
import time
import pandas as pd
import argparse
from datetime import datetime
# --- 日誌設定 ---
log_time = datetime.now().strftime('%Y%m%d_%H%M%S')
log_directory = 'logs'
os.makedirs(log_directory, exist_ok=True)
log_filename = f"balancing_{log_time}.log"
logging.basicConfig(level=logging.INFO, filename=os.path.join(log_directory, log_filename), filemode='w', format='%(asctime)s - %(levelname)s - %(message)s', encoding='utf-8')
console_handler = logging.StreamHandler(); console_handler.setLevel(logging.INFO); formatter = logging.Formatter('%(asctime)s - %(levelname)s - %(message)s'); console_handler.setFormatter(formatter); logging.getLogger().addHandler(console_handler)
def balance_dataset_by_mission(df: pd.DataFrame) -> pd.DataFrame:
"""
對資料集進行策略性平衡,確保每個 mission 的數量大致相同。
"""
if df.empty:
logging.warning("警告:傳入的 DataFrame 為空,無法進行平衡處理。")
return pd.DataFrame()
logging.info("--- 開始執行策略性數據集平衡 ---")
mission_counts = df['mission'].value_counts()
logging.info("待平衡的有效資料各任務分佈:\n%s", mission_counts)
if mission_counts.empty:
logging.warning("沒有資料可供平衡,將回傳原始 DataFrame。")
return df
target_count = mission_counts.min()
logging.info(f"目標數量 (以最少任務為基準): {target_count}")
df['combination'] = df['genmodel'] + " / " + df['judgemodel']
combination_priority = df['combination'].value_counts().reset_index()
combination_priority.columns = ['combination', 'total_count']
combination_priority = combination_priority.sort_values('total_count', ascending=True) # 由低至高排序
logging.info(f"模型組合優先級 (由低至高):\n{combination_priority.to_string()}")
balanced_dfs = []
for mission, group_df in df.groupby('mission'):
if len(group_df) <= target_count:
logging.info(f"任務 '{mission}' (數量: {len(group_df)}) 無需抽樣,全部保留。")
balanced_dfs.append(group_df)
continue
logging.info(f"任務 '{mission}' (數量: {len(group_df)}) 需要下採樣至 {target_count} 筆。")
samples_to_keep = pd.DataFrame()
for _, row in combination_priority.iterrows():
priority_combo = row['combination']
combo_in_group = group_df[group_df['combination'] == priority_combo]
if combo_in_group.empty: continue
if len(samples_to_keep) + len(combo_in_group) <= target_count:
samples_to_keep = pd.concat([samples_to_keep, combo_in_group])
logging.info(f" - 保留組合 '{priority_combo}' 的全部 {len(combo_in_group)} 筆。目前總數: {len(samples_to_keep)}")
else:
remaining_needed = target_count - len(samples_to_keep)
if remaining_needed > 0:
samples_to_add = combo_in_group.sample(n=remaining_needed, random_state=42)
samples_to_keep = pd.concat([samples_to_keep, samples_to_add])
logging.info(f" - 從組合 '{priority_combo}' 中隨機抽取 {remaining_needed} 筆以補滿。")
break
balanced_dfs.append(samples_to_keep)
final_balanced_df = pd.concat(balanced_dfs).reset_index(drop=True)
if 'combination' in final_balanced_df.columns:
final_balanced_df = final_balanced_df.drop(columns=['combination'])
logging.info("--- 數據集平衡處理完成 ---")
logging.info("平衡後有效資料各任務分佈:\n%s", final_balanced_df['mission'].value_counts())
return final_balanced_df
# --- 【核心修改】新增一個專門用來分類儲存的函式 ---
def save_balanced_classified_files(final_df: pd.DataFrame, output_dir: str):
"""
將最終的、已平衡的 DataFrame 分類儲存成多個獨立的 JSON 檔案。
"""
if final_df.empty:
logging.warning("最終 DataFrame 為空,沒有任何檔案可以儲存。")
return
os.makedirs(output_dir, exist_ok=True)
logging.info(f"開始將平衡後的分類結果儲存至 '{output_dir}' 目錄...")
# 將 DataFrame 轉回 list of dicts 以便處理
all_items_metadata = final_df.to_dict('records')
# 準備 buckets (與 process_data_v2.py 中的邏輯相同)
final_classified_data = {}
format_error_log = []
deduplicated_log = []
for item in all_items_metadata:
if item.get('status') == 'FormatError':
format_error_log.append(item)
elif item.get('status') == 'Duplicate':
deduplicated_log.append(item)
elif item.get('status') == 'Valid':
# 確保mission 欄位存在
mission = item.get('mission')
bucket_key = f"traditional_{mission}"
if bucket_key not in final_classified_data:
final_classified_data[bucket_key] = []
final_classified_data[bucket_key].append(item)
# 合併所有要輸出的資料
all_output_data = {**final_classified_data, 'format_error': format_error_log, 'deduplicated_log': deduplicated_log}
# 遍歷並儲存所有檔案
for filename_key, data in all_output_data.items():
if not data: continue
filename = os.path.join(output_dir, f"{filename_key}.json")
try:
for bucket_data in final_classified_data.values():
for item in bucket_data:
item.pop('status', None); item.pop('error_reason', None); item.pop('text_for_hash', None); item.pop('language', None); item.pop('combination', None); item.pop('internal_id', None); item.pop('compared_with_id', None); item.pop('compared_with_qid', None)
with open(filename, 'w', encoding='utf-8') as f:
json.dump(data, f, ensure_ascii=False, indent=4)
logging.info(f"成功儲存 {len(data)} 筆資料至: {filename}")
except Exception as e:
logging.error(f"儲存檔案 {filename} 時發生錯誤: {e}")
# --- 主程式執行區 ---
if __name__ == "__main__":
parser = argparse.ArgumentParser(description="從多個分類好的 JSON 檔案中讀取資料,進行平衡後再分類儲存。")
parser.add_argument(
'--input_dir',
type=str,
default='final_classified_output', # 【建議】保持目錄名稱簡潔
help='包含多個分類好的 JSON 檔案的輸入目錄。'
)
parser.add_argument(
'--output_dir',
type=str,
default='final_classified_output_new', # 【建議】修改為更清晰的名稱
help='最終輸出的分類檔案存放目錄。'
)
args = parser.parse_args()
start_time = time.time()
# --- 步驟 1: 讀取並處理「有效資料」---
target_files = ['traditional_od.json', 'traditional_press.json', 'traditional_petition.json', 'traditional_qa.json']
all_valid_data = []
logging.info(f"準備從 '{args.input_dir}' 目錄讀取指定的分類檔案...")
for filename in target_files:
file_path = os.path.join(args.input_dir, filename)
if os.path.exists(file_path):
try:
with open(file_path, 'r', encoding='utf-8') as f:
data = json.load(f)
all_valid_data.extend(data)
logging.info(f"成功從 '{filename}' 讀取 {len(data)} 筆資料。")
except Exception as e: logging.error(f"讀取檔案 {filename} 時發生錯誤: {e}")
else: logging.warning(f"警告:找不到目標檔案 {file_path},將跳過。")
if not all_valid_data:
logging.error("未能讀取到任何有效資料,程式即將終止。")
exit()
valid_df = pd.DataFrame(all_valid_data)
logging.info(f"共讀取 {len(valid_df)} 筆有效資料準備進行平衡。")
required_cols = ['mission', 'genmodel', 'judgemodel']
if not all(col in valid_df.columns for col in required_cols):
logging.error("錯誤:輸入的 JSON 檔案缺少必要的欄位 (mission, genmodel, judgemodel)。請確保第一階段的腳本沒有移除這些欄位。")
exit()
# 【修正】明確地給予 valid_df 一個 status 欄位
valid_df['status'] = 'Valid'
balanced_valid_df = balance_dataset_by_mission(valid_df)
# --- 步驟 2: 讀取並處理「日誌資料」---
all_other_data = []
# 【修正】處理巢狀結構
if os.path.exists(os.path.join(args.input_dir, 'format_error.json')):
with open(os.path.join(args.input_dir, 'format_error.json'), 'r', encoding='utf-8') as f:
for item in json.load(f):
# 提取出真正的資料,並確保它有 status
details = item.get('format_error_item_details', {})
if details:
details['status'] = 'FormatError'
all_other_data.append(details)
if os.path.exists(os.path.join(args.input_dir, 'deduplicated_log.json')):
with open(os.path.join(args.input_dir, 'deduplicated_log.json'), 'r', encoding='utf-8') as f:
for item in json.load(f):
details = item.get('removed_item_details', {})
if details:
details['status'] = 'Duplicate'
all_other_data.append(details)
other_df = pd.DataFrame(all_other_data)
logging.info(f"共讀取 {len(other_df)} 筆日誌資料 (錯誤/重複)。")
# --- 步驟 3: 合併並儲存 ---
final_df = pd.concat([balanced_valid_df, other_df], ignore_index=True)
logging.info(f"資料重新組合完成,最終總筆數為: {len(final_df)}")
save_balanced_classified_files(final_df, args.output_dir)
duration = time.time() - start_time
logging.info(f"平衡腳本執行完畢,共耗時 {duration:.2f} 秒。")
print(f"\n平衡腳本執行完畢,共耗時 {duration:.2f} 秒。")
print(f"所有平衡後的分類檔案已儲存至 '{os.path.abspath(args.output_dir)}' 目錄。")