-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathdataloader.py
More file actions
123 lines (107 loc) · 3.96 KB
/
Copy pathdataloader.py
File metadata and controls
123 lines (107 loc) · 3.96 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
import os
import json
import random
import numpy as np
from PIL import Image
import pandas as pd
from sklearn.model_selection import train_test_split
import warnings
warnings.filterwarnings("ignore")
def read_scores(path):
with open(path) as f:
return json.load(f)
def find_image(image_dir, fname):
path = os.path.join(image_dir, fname)
if os.path.exists(path): return path
base, _ = os.path.splitext(fname)
for ext in ['.tif', '.tiff', '.png', '.jpg']:
alt = os.path.join(image_dir, base + ext)
if os.path.exists(alt): return alt
return None
def load_data(image_dir, score_dict):
data = []
for fname, score in score_dict.items():
p = find_image(image_dir, fname)
if not p:
print(f"Missing: {fname}")
continue
img = np.array(Image.open(p))
if img.ndim == 3: img = img[..., 0]
data.append((img, score, os.path.basename(p)))
return data
def load_kadid10k_data(kadid_dir):
csv_path = os.path.join(kadid_dir, "dmos.csv")
images_dir = os.path.join(kadid_dir, "images")
df = pd.read_csv(csv_path, sep=',')
data = []
for idx, row in df.iterrows():
img_name = row['dist_img']
dmos = float(row['dmos']) - 1.0 # [1,5] → [0,4]
img_path = os.path.join(images_dir, img_name)
if os.path.exists(img_path):
img = Image.open(img_path)
if img.mode == 'RGB':
img = img.convert('L') # Always grayscale!
img = img.resize((512, 512), Image.BILINEAR) # <<--- ADD THIS LINE
img = np.array(img)
data.append((img, dmos, img_name))
else:
print(f"Missing: {img_path}")
print(f"KADID10k: Loaded {len(data)} images")
return data
# def stratified_split(data, ratio=0.9, n_bins=10, seed=42):
# """
# Stratified split based on scores.
# Args:
# data: list of (img, score, fname)
# ratio: proportion for train
# n_bins: number of bins for discretizing continuous scores
# """
# scores = np.array([d[1] for d in data])
# # Bin continuous scores into categories for stratification
# bins = np.linspace(scores.min(), scores.max(), n_bins)
# y_binned = np.digitize(scores, bins, right=True)
# train, val = train_test_split(
# data,
# test_size=1-ratio,
# stratify=y_binned,
# random_state=seed
# )
# return train, val
# def split(data, ratio=0.9):
# random.shuffle(data)
# i = int(len(data) * ratio)
# return data[:i], data[i:]
def split(data, ratio=0.9, n_bins=10, seed=42):
"""
Stratified split based on IQA scores to ensure equal distribution.
Args:
data: list of (img, score, fname)
ratio: proportion for training set (default 0.9)
n_bins: number of bins for discretizing continuous scores
seed: random seed
"""
scores = np.array([d[1] for d in data]) # Extract IQA scores
# Bin scores into discrete categories for stratification
bins = np.linspace(scores.min(), scores.max(), n_bins)
y_binned = np.digitize(scores, bins, right=True)
train, val = train_test_split(
data,
test_size=1 - ratio,
stratify=y_binned,
random_state=seed,
)
print(f"[split] Stratified: train={len(train)} val={len(val)} bins={n_bins}")
return train, val
if __name__ == "__main__":
random.seed(42)
base = os.path.dirname(os.path.abspath(__file__))
train_dir = os.path.join(base, "dataset/ldctiqac/train/image")
train_json = os.path.join(base, "dataset/ldctiqac/train/train.json")
test_dir = os.path.join(base, "dataset/ldctiqac/test/images")
test_json = os.path.join(base, "dataset/ldctiqac/test/test.json")
train_full = load_data(train_dir, read_scores(train_json))
train, val = split(train_full)
print(f"Train: {len(train)}, Val: {len(val)}")
test = load_data(test_dir, read_scores(test_json))
print(f"Test: {len(test)}\nData Loaded!")