chris-123 commited on
Commit
11d7bc2
·
verified ·
1 Parent(s): a229974

Upload ensemble_results.py

Browse files
Files changed (1) hide show
  1. ensemble_results.py +407 -0
ensemble_results.py ADDED
@@ -0,0 +1,407 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import re
3
+ import json
4
+ import glob
5
+ import copy
6
+ from collections import Counter, defaultdict
7
+ from statistics import median
8
+
9
+ CRITERIA = [
10
+ "Color Harmony",
11
+ "Visual Style Consistency",
12
+ "Sharpness",
13
+ "Light and Shadow Modeling",
14
+ "Creativity and Originality",
15
+ "Exposure Control",
16
+ "Application of Classical Composition Principles",
17
+ "Depth of Field and Layering",
18
+ "Visual Center Stability",
19
+ "Visual Flow Guidance",
20
+ "Structural Support Stability",
21
+ "Appropriateness of Negative Space",
22
+ "Subject Integrity",
23
+ ]
24
+
25
+ # LEVEL_ORDER = {"Poor": 0, "Medium": 1, "Good": 2}
26
+ # LEVEL_INV = {0: "Poor", 1: "Medium", 2: "Good"}
27
+
28
+
29
+ LEVEL_ORDER = {"Poor": 0, "Medium": 1, "Good": 2}
30
+ LEVEL_INV = {0: "A", 1: "B", 2: "C"}
31
+
32
+
33
+ def normalize_level(x):
34
+ if not isinstance(x, str):
35
+ return None
36
+ x = x.strip().lower()
37
+ mp = {
38
+ "poor": "Poor",
39
+ "medium": "Medium",
40
+ "good": "Good",
41
+ }
42
+ return mp.get(x)
43
+
44
+
45
+ def basename_from_item(item):
46
+ img_path = item.get("images", [{}])[0].get("path", "")
47
+ return os.path.basename(img_path)
48
+
49
+
50
+ def parse_response_raw(resp):
51
+ """
52
+ 支持:
53
+ - "{\"total_score\": 84}"
54
+ - "41"
55
+ - "{\"criteria\": {...}}"
56
+ - "[\"Medium\", ...]"
57
+ - "{\"answer\": \"C\"}"
58
+ """
59
+ if isinstance(resp, (dict, list, int, float)):
60
+ return resp
61
+
62
+ if not isinstance(resp, str):
63
+ return None
64
+
65
+ s = resp.strip()
66
+
67
+ # 纯数字分数
68
+ if re.fullmatch(r"-?\d+(\.\d+)?", s):
69
+ return float(s)
70
+
71
+ # 标准 JSON
72
+ try:
73
+ return json.loads(s)
74
+ except Exception:
75
+ pass
76
+
77
+ # 兜底:提取 {...}
78
+ m = re.search(r"\{.*\}", s, flags=re.S)
79
+ if m:
80
+ try:
81
+ return json.loads(m.group(0))
82
+ except Exception:
83
+ pass
84
+
85
+ # 兜底:提取 [...]
86
+ m = re.search(r"\[.*\]", s, flags=re.S)
87
+ if m:
88
+ try:
89
+ return json.loads(m.group(0))
90
+ except Exception:
91
+ pass
92
+
93
+ return None
94
+
95
+
96
+ def iter_json_or_jsonl(path):
97
+ with open(path, "r", encoding="utf-8") as f:
98
+ text = f.read().strip()
99
+
100
+ if not text:
101
+ return []
102
+
103
+ try:
104
+ obj = json.loads(text)
105
+ if isinstance(obj, list):
106
+ return obj
107
+ if isinstance(obj, dict):
108
+ return [obj]
109
+ except Exception:
110
+ pass
111
+
112
+ rows = []
113
+ for line in text.splitlines():
114
+ line = line.strip()
115
+ if line:
116
+ rows.append(json.loads(line))
117
+ return rows
118
+
119
+
120
+ def read_all_files(folder, recursive=True):
121
+ pattern = "**/*.json*" if recursive else "*.json*"
122
+ files = sorted(glob.glob(os.path.join(folder, pattern), recursive=recursive))
123
+
124
+ rows = []
125
+ for fp in files:
126
+ try:
127
+ rows.extend(iter_json_or_jsonl(fp))
128
+ except Exception as e:
129
+ print(f"[WARN] failed to read {fp}: {e}")
130
+ return rows
131
+
132
+
133
+ def parse_score_folder(folder):
134
+ pred = defaultdict(list)
135
+
136
+ for item in read_all_files(folder):
137
+ name = basename_from_item(item)
138
+ data = parse_response_raw(item.get("response", ""))
139
+
140
+ score = None
141
+
142
+ if isinstance(data, dict):
143
+ score = data.get("total_score")
144
+ elif isinstance(data, (int, float)):
145
+ score = data
146
+ elif isinstance(data, str):
147
+ if re.fullmatch(r"-?\d+(\.\d+)?", data.strip()):
148
+ score = float(data.strip())
149
+
150
+ if name and score is not None:
151
+ try:
152
+ score = float(score)
153
+ score = max(0, min(100, score))
154
+ pred[name].append(score)
155
+ except Exception:
156
+ pass
157
+
158
+ return pred
159
+
160
+
161
+ def parse_level_folder(folder):
162
+ pred = defaultdict(lambda: defaultdict(list))
163
+
164
+ for item in read_all_files(folder):
165
+ name = basename_from_item(item)
166
+ data = parse_response_raw(item.get("response", ""))
167
+
168
+ if not name:
169
+ continue
170
+
171
+ # 格式1:{"criteria": {"Color Harmony": "Good", ...}}
172
+ if isinstance(data, dict) and isinstance(data.get("criteria"), dict):
173
+ criteria = data["criteria"]
174
+ for c in CRITERIA:
175
+ lv = normalize_level(criteria.get(c))
176
+ if lv:
177
+ pred[name][c].append(lv)
178
+
179
+ # 格式2:["Medium", "Medium", ..., 共13个]
180
+ elif isinstance(data, list):
181
+ for c, lv_raw in zip(CRITERIA, data):
182
+ lv = normalize_level(lv_raw)
183
+ if lv:
184
+ pred[name][c].append(lv)
185
+
186
+ return pred
187
+
188
+
189
+ def parse_reason_folder(folder):
190
+ pred = defaultdict(list)
191
+
192
+ for item in read_all_files(folder):
193
+ name = basename_from_item(item)
194
+ data = parse_response_raw(item.get("response", ""))
195
+
196
+ ans = None
197
+
198
+ if isinstance(data, dict):
199
+ ans = data.get("answer")
200
+ elif isinstance(data, str):
201
+ ans = data
202
+
203
+ if name and isinstance(ans, str):
204
+ ans = ans.strip().upper()
205
+ if ans in {"A", "B", "C", "D"}:
206
+ pred[name].append(ans)
207
+
208
+ return pred
209
+
210
+
211
+ def majority_vote(values, default=None):
212
+ values = [v for v in values if v is not None]
213
+ if not values:
214
+ return default
215
+
216
+ cnt = Counter(values)
217
+
218
+ # 平票时按第一次出现顺序
219
+ return max(cnt.keys(), key=lambda x: (cnt[x], -values.index(x)))
220
+
221
+
222
+ def ensemble_scores(score_dicts, method="mean"):
223
+ merged = defaultdict(list)
224
+
225
+ # print("merged is", merged)
226
+
227
+ for d in score_dicts:
228
+ for name, scores in d.items():
229
+ merged[name].extend(scores)
230
+
231
+ out = {}
232
+ for name, scores in merged.items():
233
+
234
+
235
+ if method == "median":
236
+ val = median(scores)
237
+ else:
238
+ val = sum(scores) / len(scores)
239
+
240
+ # print("scores are", name, scores, val)
241
+
242
+ out[name] = int(round(max(0, min(100, val))))
243
+ # out[name] = int(round(val))
244
+
245
+ return out
246
+
247
+
248
+
249
+ #####################
250
+
251
+ LEVEL_SCORE = {
252
+ "Poor": 2.5,
253
+ "Medium": 6.0,
254
+ "Good": 8.5,
255
+ }
256
+
257
+
258
+ def score_to_level(score):
259
+ if 0 <= score < 5:
260
+ return "A"
261
+ elif 5 <= score < 7:
262
+ return "B"
263
+ elif 7 <= score <= 10:
264
+ return "C"
265
+ else:
266
+ # 兜底,防止异常值
267
+ score = max(0, min(10, score))
268
+ if score < 5:
269
+ return "A"
270
+ elif score < 7:
271
+ return "B"
272
+ return "C"
273
+
274
+
275
+ def ensemble_levels(level_dicts, method="score_mean"):
276
+ merged = defaultdict(lambda: defaultdict(list))
277
+
278
+ for d in level_dicts:
279
+ for name, cd in d.items():
280
+ for c, levels in cd.items():
281
+ merged[name][c].extend(levels)
282
+
283
+ out = defaultdict(dict)
284
+
285
+ for name, cd in merged.items():
286
+ for c in CRITERIA:
287
+ vals = cd.get(c, [])
288
+ if not vals:
289
+ continue
290
+
291
+ if method == "score_mean":
292
+ nums = [LEVEL_SCORE[v] for v in vals if v in LEVEL_SCORE]
293
+ if nums:
294
+ avg_score = sum(nums) / len(nums)
295
+ out[name][c] = score_to_level(avg_score)
296
+
297
+ elif method == "vote":
298
+ out[name][c] = majority_vote(vals, default="Medium")
299
+
300
+ elif method == "ordinal_mean":
301
+ nums = [LEVEL_ORDER[v] for v in vals if v in LEVEL_ORDER]
302
+ if nums:
303
+ out[name][c] = LEVEL_INV[int(round(sum(nums) / len(nums)))]
304
+
305
+ return out
306
+
307
+
308
+
309
+ ######################
310
+
311
+
312
+ def ensemble_answers(reason_dicts):
313
+ merged = defaultdict(list)
314
+
315
+ for d in reason_dicts:
316
+ for name, answers in d.items():
317
+ merged[name].extend(answers)
318
+
319
+ return {
320
+ name: majority_vote(answers, default="A")
321
+ for name, answers in merged.items()
322
+ }
323
+
324
+
325
+ def build_submission(
326
+ template_path,
327
+ score_model_folders,
328
+ level_model_folders,
329
+ reason_model_folders,
330
+ output_path,
331
+ score_method="mean",
332
+ level_method="vote",
333
+ ):
334
+ score_dicts = [parse_score_folder(p) for p in score_model_folders]
335
+ level_dicts = [parse_level_folder(p) for p in level_model_folders]
336
+ reason_dicts = [parse_reason_folder(p) for p in reason_model_folders]
337
+
338
+ score_ens = ensemble_scores(score_dicts, method=score_method)
339
+ level_ens = ensemble_levels(level_dicts, method=level_method)
340
+ answer_ens = ensemble_answers(reason_dicts)
341
+
342
+ with open(template_path, "r", encoding="utf-8") as f:
343
+ result = json.load(f)
344
+
345
+ missing_score = 0
346
+ missing_level = 0
347
+ missing_answer = 0
348
+
349
+ for item in result:
350
+ name = item["image_path"]
351
+
352
+ if name in score_ens:
353
+ item["total_score"] = score_ens[name]
354
+ else:
355
+ missing_score += 1
356
+
357
+ for c in CRITERIA:
358
+ if name in level_ens and c in level_ens[name]:
359
+ item["criteria"][c]["level"] = level_ens[name][c]
360
+ else:
361
+ missing_level += 1
362
+
363
+ if name in answer_ens:
364
+ item["answer"] = answer_ens[name]
365
+ else:
366
+ missing_answer += 1
367
+
368
+ with open(output_path, "w", encoding="utf-8") as f:
369
+ json.dump(result, f, ensure_ascii=False, indent=2)
370
+
371
+ print(f"Saved to: {output_path}")
372
+ print(f"Images: {len(result)}")
373
+ print(f"Missing score images: {missing_score}")
374
+ print(f"Missing level fields: {missing_level}")
375
+ print(f"Missing answer images: {missing_answer}")
376
+
377
+
378
+ if __name__ == "__main__":
379
+
380
+ # from pathlib import Path
381
+ # ROOT = Path("/mnt/shared-storage-user/zhuxiaorong/liyunhao_data/my_code/cvpr26_challenge")
382
+
383
+
384
+ TEMPLATE_PATH = "track_1_test_demo.json"
385
+ OUTPUT_PATH = "./track_1_test.json"
386
+
387
+
388
+ SCORE_MODEL_FOLDERS = [
389
+ "./result-score"
390
+ ]
391
+ LEVEL_MODEL_FOLDERS = [
392
+ "./result-level"
393
+ ]
394
+ REASON_MODEL_FOLDERS = [
395
+ "./result-reason"
396
+ ]
397
+
398
+
399
+ build_submission(
400
+ template_path=TEMPLATE_PATH,
401
+ score_model_folders=SCORE_MODEL_FOLDERS,
402
+ level_model_folders=LEVEL_MODEL_FOLDERS,
403
+ reason_model_folders=REASON_MODEL_FOLDERS,
404
+ output_path=OUTPUT_PATH,
405
+ score_method="mean",
406
+ level_method="score_mean",
407
+ )