vislearnlabpy 0.0.2.2__tar.gz → 0.0.2.4__tar.gz
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- {vislearnlabpy-0.0.2.2/src/vislearnlabpy.egg-info → vislearnlabpy-0.0.2.4}/PKG-INFO +1 -1
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/pyproject.toml +1 -1
- vislearnlabpy-0.0.2.4/src/vislearnlabpy/embeddings/stimuli_loader.py +377 -0
- vislearnlabpy-0.0.2.4/src/vislearnlabpy/models/dinov3_model.py +20 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/models/multimodal_model.py +0 -39
- vislearnlabpy-0.0.2.4/src/vislearnlabpy/models/vision_model.py +54 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4/src/vislearnlabpy.egg-info}/PKG-INFO +1 -1
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy.egg-info/SOURCES.txt +3 -1
- vislearnlabpy-0.0.2.2/src/vislearnlabpy/embeddings/stimuli_loader.py +0 -207
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/LICENSE +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/README.md +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/setup.cfg +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/__init__.py +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/drawings/drawing.py +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/drawings/svg_render_helpers.py +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/embeddings/embedding_store.py +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/embeddings/generate_embeddings.py +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/embeddings/similarity_generator.py +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/embeddings/similarity_utils.py +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/embeddings/utils.py +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/extractions/drawingtask_extractor.py +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/extractions/mongo_extractor.py +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/models/clip_model.py +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/models/feature_generator.py +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy.egg-info/dependency_links.txt +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy.egg-info/requires.txt +0 -0
- {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy.egg-info/top_level.txt +0 -0
|
@@ -0,0 +1,377 @@
|
|
|
1
|
+
from dataclasses import dataclass
|
|
2
|
+
import os
|
|
3
|
+
import pandas as pd
|
|
4
|
+
import re
|
|
5
|
+
from PIL import Image
|
|
6
|
+
from torch.utils.data import Dataset, DataLoader
|
|
7
|
+
from vislearnlabpy.embeddings.utils import process_csv
|
|
8
|
+
import random
|
|
9
|
+
from typing import Callable, Any
|
|
10
|
+
import warnings
|
|
11
|
+
from torchvision import transforms
|
|
12
|
+
import numpy as np
|
|
13
|
+
from scipy import ndimage
|
|
14
|
+
from skimage.morphology import skeletonize, binary_dilation, disk
|
|
15
|
+
from skimage import img_as_bool
|
|
16
|
+
|
|
17
|
+
random.seed(2)
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
@dataclass
|
|
21
|
+
class ImgExtractionSettings:
|
|
22
|
+
resize_dim: int = 256
|
|
23
|
+
crop_dim: int = 224
|
|
24
|
+
apply_content_crop: bool = True
|
|
25
|
+
apply_center_crop: bool = False
|
|
26
|
+
use_thumbnail: bool = False
|
|
27
|
+
change_stroke_color: bool = False
|
|
28
|
+
stroke_color: tuple = (0, 0, 0)
|
|
29
|
+
stroke_threshold: int = 200
|
|
30
|
+
bg_threshold: int = 200 # threshold to consider a pixel as background (0-255)
|
|
31
|
+
bg_component_size: int = 10 # minimum size of connected components to keep (in pixels)
|
|
32
|
+
filter_edge_artifacts: bool = False # whether to filter out edge artifacts during cropping
|
|
33
|
+
normalize_stroke_thickness: bool = False
|
|
34
|
+
stroke_target_thickness: int = 2 # target thickness for stroke normalization
|
|
35
|
+
double_resize: bool = False # whether to resize twice (once with thumbnail, once after with resize func)
|
|
36
|
+
|
|
37
|
+
class ImageExtractor:
|
|
38
|
+
@staticmethod
|
|
39
|
+
def RGBA2RGB(img, background_color=(255, 255, 255)):
|
|
40
|
+
"""Alpha composite an RGBA Image with a specified color.
|
|
41
|
+
Source: http://stackoverflow.com/a/9459208/284318
|
|
42
|
+
|
|
43
|
+
Keyword Arguments:
|
|
44
|
+
img -- PIL RGBA Image object
|
|
45
|
+
background_color -- Tuple r, g, b (default 255, 255, 255)
|
|
46
|
+
"""
|
|
47
|
+
if img.mode in ('RGBA', 'LA'): # Image has transparency
|
|
48
|
+
# Create white background
|
|
49
|
+
white_bg = Image.new('RGB', img.size, background_color)
|
|
50
|
+
# Paste the image onto white background
|
|
51
|
+
# The image itself acts as the alpha mask
|
|
52
|
+
white_bg.paste(img, mask=img.split()[-1] if img.mode == 'RGBA' else None)
|
|
53
|
+
img = white_bg
|
|
54
|
+
elif img.mode == 'P': # Palette mode, might have transparency
|
|
55
|
+
img = img.convert('RGBA')
|
|
56
|
+
white_bg = Image.new('RGB', img.size, background_color)
|
|
57
|
+
white_bg.paste(img, mask=img.split()[-1])
|
|
58
|
+
img = white_bg
|
|
59
|
+
else:
|
|
60
|
+
# Image is already RGB/L, convert to RGB if needed
|
|
61
|
+
img = img.convert('RGB')
|
|
62
|
+
return img
|
|
63
|
+
|
|
64
|
+
@staticmethod
|
|
65
|
+
def change_stroke_color(img, target_color=(0, 0, 0), threshold=200):
|
|
66
|
+
"""
|
|
67
|
+
Change the color of strokes/dark pixels in an image.
|
|
68
|
+
|
|
69
|
+
Args:
|
|
70
|
+
img -- PIL Image object
|
|
71
|
+
target_color -- Tuple (r, g, b) for the new stroke color
|
|
72
|
+
threshold -- Brightness threshold to identify strokes (0-255)
|
|
73
|
+
Pixels darker than this are considered strokes
|
|
74
|
+
|
|
75
|
+
Returns:
|
|
76
|
+
PIL Image with recolored strokes
|
|
77
|
+
"""
|
|
78
|
+
# Convert to RGBA if not already
|
|
79
|
+
if img.mode != 'RGBA':
|
|
80
|
+
img = img.convert('RGBA')
|
|
81
|
+
|
|
82
|
+
# Convert to numpy array
|
|
83
|
+
img_array = np.array(img)
|
|
84
|
+
|
|
85
|
+
# Calculate brightness (average of RGB channels)
|
|
86
|
+
brightness = img_array[:, :, :3].mean(axis=2)
|
|
87
|
+
|
|
88
|
+
# Create mask for stroke pixels (darker than threshold)
|
|
89
|
+
stroke_mask = brightness < threshold
|
|
90
|
+
|
|
91
|
+
# Apply new color to stroke pixels while preserving alpha
|
|
92
|
+
result = img_array.copy()
|
|
93
|
+
result[stroke_mask, 0] = target_color[0] # R
|
|
94
|
+
result[stroke_mask, 1] = target_color[1] # G
|
|
95
|
+
result[stroke_mask, 2] = target_color[2] # B
|
|
96
|
+
|
|
97
|
+
return Image.fromarray(result)
|
|
98
|
+
|
|
99
|
+
def normalize_stroke_thickness(img, target_thickness=2):
|
|
100
|
+
# Binarize
|
|
101
|
+
arr = np.array(img.convert('L'))
|
|
102
|
+
binary = arr < 128 # assuming dark strokes on light background
|
|
103
|
+
|
|
104
|
+
# Skeletonize to 1-pixel lines
|
|
105
|
+
skeleton = skeletonize(binary)
|
|
106
|
+
|
|
107
|
+
# Dilate to consistent thickness
|
|
108
|
+
if target_thickness > 1:
|
|
109
|
+
skeleton = binary_dilation(skeleton, disk(target_thickness//2))
|
|
110
|
+
|
|
111
|
+
# Convert back to image (white background, black strokes)
|
|
112
|
+
result = (~skeleton * 255).astype(np.uint8)
|
|
113
|
+
return Image.fromarray(result)
|
|
114
|
+
|
|
115
|
+
@staticmethod
|
|
116
|
+
def _filter_edge_artifacts(large_components, labeled, component_sizes, img_width, img_height):
|
|
117
|
+
for comp_id in range(1, len(component_sizes)):
|
|
118
|
+
if large_components[comp_id]:
|
|
119
|
+
comp_mask = (labeled == comp_id)
|
|
120
|
+
rows, cols = np.where(comp_mask)
|
|
121
|
+
|
|
122
|
+
if len(rows) > 0 or len(cols) > 0:
|
|
123
|
+
ylb, yub = min(rows), max(rows)
|
|
124
|
+
xlb, xub = min(cols), max(cols)
|
|
125
|
+
width = xub - xlb + 1
|
|
126
|
+
height = yub - ylb + 1
|
|
127
|
+
x_tolerance = 10
|
|
128
|
+
y_tolerance = 100
|
|
129
|
+
touches_top = ylb <= y_tolerance
|
|
130
|
+
touches_bottom = yub >= img_height - 1 - y_tolerance
|
|
131
|
+
touches_left = xlb <= x_tolerance
|
|
132
|
+
touches_right = xub >= img_width - 1 - x_tolerance
|
|
133
|
+
# Check if it's a thin strip spanning nearly the entire edge
|
|
134
|
+
# vertical strip is more likely to not be image artifact
|
|
135
|
+
is_thin_vertical_strip = (touches_top or touches_bottom and width < img_width * 0.2)
|
|
136
|
+
is_thin_horizontal_strip = (touches_left and touches_right and height < img_height * 0.2)
|
|
137
|
+
if is_thin_vertical_strip or is_thin_horizontal_strip or (width / height >= 10) or (height / width >= 10):
|
|
138
|
+
large_components[comp_id] = False
|
|
139
|
+
return large_components
|
|
140
|
+
|
|
141
|
+
@staticmethod
|
|
142
|
+
def crop_to_content(img, apply_content_crop, stroke_threshold=200, min_component_size=10, filter_edge_artifacts=False):
|
|
143
|
+
"""Crop image to remove white space around content, ignoring small artifacts."""
|
|
144
|
+
if not apply_content_crop:
|
|
145
|
+
return img
|
|
146
|
+
|
|
147
|
+
arr = np.asarray(img)
|
|
148
|
+
img_width, img_height = img.size
|
|
149
|
+
# Create binary mask of non-white pixels
|
|
150
|
+
mask = np.all(arr < stroke_threshold, axis=2) if arr.ndim == 3 else arr < stroke_threshold
|
|
151
|
+
|
|
152
|
+
# Label connected components
|
|
153
|
+
labeled, num_features = ndimage.label(mask)
|
|
154
|
+
|
|
155
|
+
# Find component sizes
|
|
156
|
+
component_sizes = np.bincount(labeled.ravel())
|
|
157
|
+
# Component 0 is background, so start from 1
|
|
158
|
+
large_components = component_sizes >= min_component_size
|
|
159
|
+
large_components[0] = False # Don't include background
|
|
160
|
+
# Filter out edge artifacts (components that span entire width or height)
|
|
161
|
+
if filter_edge_artifacts:
|
|
162
|
+
large_components = ImageExtractor._filter_edge_artifacts(large_components, labeled, component_sizes, img_width, img_height)
|
|
163
|
+
|
|
164
|
+
# Create cleaned mask with only large components
|
|
165
|
+
cleaned_mask = large_components[labeled]
|
|
166
|
+
if arr.ndim == 3:
|
|
167
|
+
white_img = np.ones_like(arr) * 255
|
|
168
|
+
else:
|
|
169
|
+
white_img = np.ones_like(arr) * 255
|
|
170
|
+
|
|
171
|
+
# Copy only the content pixels (keep background white)
|
|
172
|
+
white_img[cleaned_mask] = arr[cleaned_mask]
|
|
173
|
+
img = Image.fromarray(white_img.astype(np.uint8))
|
|
174
|
+
try:
|
|
175
|
+
rows, cols = np.where(cleaned_mask)
|
|
176
|
+
ylb = min(rows) # top bound
|
|
177
|
+
yub = max(rows) # bottom bound
|
|
178
|
+
xlb = min(cols) # left bound
|
|
179
|
+
xub = max(cols) # right bound
|
|
180
|
+
width = xub - xlb
|
|
181
|
+
height = yub - ylb
|
|
182
|
+
max_dim = max(width, height)
|
|
183
|
+
|
|
184
|
+
x_center = (xlb + xub) // 2
|
|
185
|
+
y_center = (ylb + yub) // 2
|
|
186
|
+
|
|
187
|
+
lb_x = x_center - max_dim // 2
|
|
188
|
+
ub_x = x_center + max_dim // 2
|
|
189
|
+
lb_y = y_center - max_dim // 2
|
|
190
|
+
ub_y = y_center + max_dim // 2
|
|
191
|
+
|
|
192
|
+
# Shift if out of bounds (to maintain square)
|
|
193
|
+
if lb_x < 0:
|
|
194
|
+
ub_x -= lb_x
|
|
195
|
+
lb_x = 0
|
|
196
|
+
if lb_y < 0:
|
|
197
|
+
ub_y -= lb_y
|
|
198
|
+
lb_y = 0
|
|
199
|
+
if ub_x > img.size[0]:
|
|
200
|
+
lb_x -= (ub_x - img.size[0])
|
|
201
|
+
ub_x = img.size[0]
|
|
202
|
+
if ub_y > img.size[1]:
|
|
203
|
+
lb_y -= (ub_y - img.size[1])
|
|
204
|
+
ub_y = img.size[1]
|
|
205
|
+
|
|
206
|
+
img = img.crop((lb_x, lb_y, ub_x, ub_y))
|
|
207
|
+
except ValueError:
|
|
208
|
+
print('Blank image - skipping crop')
|
|
209
|
+
return img
|
|
210
|
+
|
|
211
|
+
@staticmethod
|
|
212
|
+
def get_transformations(settings: ImgExtractionSettings=ImgExtractionSettings()):
|
|
213
|
+
"""Load image transformations for dataloader.
|
|
214
|
+
|
|
215
|
+
Args:
|
|
216
|
+
settings: ImgExtractionSettings dataclass instance
|
|
217
|
+
"""
|
|
218
|
+
|
|
219
|
+
def combined_transform(image):
|
|
220
|
+
# Step 1: Apply thumbnail resizing if requested
|
|
221
|
+
if settings.use_thumbnail:
|
|
222
|
+
image.thumbnail((settings.resize_dim, settings.resize_dim), Image.Resampling.LANCZOS)
|
|
223
|
+
image.thumbnail((settings.crop_dim, settings.crop_dim), Image.Resampling.LANCZOS)
|
|
224
|
+
# Step 2: Convert RGBA to RGB
|
|
225
|
+
img_rgb = ImageExtractor.RGBA2RGB(image)
|
|
226
|
+
|
|
227
|
+
# Step 3: Normalize stroke thickness if enabled
|
|
228
|
+
if settings.normalize_stroke_thickness:
|
|
229
|
+
img_rgb = ImageExtractor.normalize_stroke_thickness(
|
|
230
|
+
img_rgb,
|
|
231
|
+
settings.stroke_target_thickness
|
|
232
|
+
)
|
|
233
|
+
# change stroke color if enabled
|
|
234
|
+
if settings.change_stroke_color:
|
|
235
|
+
img_rgb = ImageExtractor.change_stroke_color(
|
|
236
|
+
img_rgb,
|
|
237
|
+
settings.stroke_color,
|
|
238
|
+
settings.stroke_threshold
|
|
239
|
+
)
|
|
240
|
+
# Step 4: Crop to content (remove whitespace) if enabled
|
|
241
|
+
img_cropped = ImageExtractor.crop_to_content(img_rgb, settings.apply_content_crop, settings.bg_threshold, settings.bg_component_size, settings.filter_edge_artifacts)
|
|
242
|
+
return img_cropped
|
|
243
|
+
|
|
244
|
+
# Build transformation pipeline
|
|
245
|
+
transform_list = []
|
|
246
|
+
if settings.apply_center_crop:
|
|
247
|
+
transform_list.append(transforms.CenterCrop(settings.crop_dim))
|
|
248
|
+
transform_list.append(transforms.Lambda(combined_transform))
|
|
249
|
+
|
|
250
|
+
# Only add regular resize if not using thumbnail
|
|
251
|
+
if not settings.use_thumbnail or settings.double_resize:
|
|
252
|
+
# resize last to ensure highest res
|
|
253
|
+
transform_list.append(transforms.Resize(settings.resize_dim))
|
|
254
|
+
|
|
255
|
+
# Add any additional transforms you might need
|
|
256
|
+
#transform_list.extend([
|
|
257
|
+
#transforms.ToTensor(),
|
|
258
|
+
# transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) # ImageNet normalization
|
|
259
|
+
#])
|
|
260
|
+
|
|
261
|
+
return transforms.Compose(transform_list)
|
|
262
|
+
|
|
263
|
+
@staticmethod
|
|
264
|
+
def save_transformed(original_file, new_file, settings: ImgExtractionSettings = None, transform=None):
|
|
265
|
+
"""
|
|
266
|
+
Load, transform, and save an image.
|
|
267
|
+
|
|
268
|
+
Args:
|
|
269
|
+
original_file: Path to input image
|
|
270
|
+
new_file: Path to save output image
|
|
271
|
+
settings: ImgExtractionSettings dataclass instance (if None, uses default)
|
|
272
|
+
transform: Custom transform (if provided, overrides settings)
|
|
273
|
+
"""
|
|
274
|
+
if settings is None:
|
|
275
|
+
settings = ImgExtractionSettings()
|
|
276
|
+
|
|
277
|
+
if transform is None:
|
|
278
|
+
transform = ImageExtractor.get_transformations(settings)
|
|
279
|
+
|
|
280
|
+
image_data = Image.open(original_file)
|
|
281
|
+
img = transform(image_data)
|
|
282
|
+
img.save(new_file)
|
|
283
|
+
|
|
284
|
+
class StimuliDataset(Dataset):
|
|
285
|
+
def __init__(self, manifest, images_folder=None, id_column=None, transform=None):
|
|
286
|
+
self.manifest = manifest
|
|
287
|
+
self.images_folder = images_folder
|
|
288
|
+
self.num_text_cols = sum(1 for c in self.manifest.columns if re.compile("text[0-9]").match(c))
|
|
289
|
+
self.num_image_cols = len([c for c in self.manifest.columns if re.compile("image[0-9]").match(c)])
|
|
290
|
+
self.id_column = id_column
|
|
291
|
+
self.transform = transform
|
|
292
|
+
|
|
293
|
+
def __len__(self):
|
|
294
|
+
return len(self.manifest)
|
|
295
|
+
|
|
296
|
+
def __getitem__(self, idx):
|
|
297
|
+
row = self.manifest.iloc[idx]
|
|
298
|
+
texts = [str(row[f"text{i}"]) for i in range(1, self.num_text_cols + 1)]
|
|
299
|
+
images = []
|
|
300
|
+
image_paths = self._get_image_paths(row)
|
|
301
|
+
full_image_paths = []
|
|
302
|
+
|
|
303
|
+
for image_path in image_paths:
|
|
304
|
+
if not os.path.isabs(image_path):
|
|
305
|
+
image_path = os.path.join(os.getcwd(), image_path)
|
|
306
|
+
if image_path is None or not os.path.exists(image_path):
|
|
307
|
+
full_image_paths.append(None)
|
|
308
|
+
images.append(None)
|
|
309
|
+
else:
|
|
310
|
+
with Image.open(image_path) as img:
|
|
311
|
+
img_copy = img.copy() # Copy the image data to memory
|
|
312
|
+
try:
|
|
313
|
+
if self.transform is not None:
|
|
314
|
+
img_copy = self.transform(img_copy)
|
|
315
|
+
else:
|
|
316
|
+
img_copy = img_copy.convert('RGB')
|
|
317
|
+
except Exception as e:
|
|
318
|
+
print(f"Warning: failed to process image {idx} ({e}); skipping transform.")
|
|
319
|
+
img_copy = img_copy.convert('RGB')
|
|
320
|
+
full_image_paths.append(image_path)
|
|
321
|
+
images.append(img_copy)
|
|
322
|
+
|
|
323
|
+
# just assuming the image path as an item id for now if there is more than one item per row.
|
|
324
|
+
# TODO: allow for the provision of both item ids and row ids
|
|
325
|
+
row_id = [row[self.id_column] if self.id_column is not None else random.randint(0, 1000000)]
|
|
326
|
+
return {"images": images, "text": texts, "item_id": full_image_paths if len(images) > 0 else row_id, "row_id": row_id}
|
|
327
|
+
|
|
328
|
+
def _get_image_paths(self, row):
|
|
329
|
+
# If this is a text only dataset
|
|
330
|
+
if self.images_folder is None and self.num_image_cols == 0:
|
|
331
|
+
return [None] * self.num_text_cols
|
|
332
|
+
# If images are not in the manifest
|
|
333
|
+
elif self.num_image_cols == 0:
|
|
334
|
+
return [os.path.join(self.images_folder, row[f"text{i}"] + ".jpg") for i in range(1, self.num_text_cols + 1)]
|
|
335
|
+
else:
|
|
336
|
+
# allowing for a secondary image path to be provided with different images in the same stimuli sets stored in different subpaths
|
|
337
|
+
return [
|
|
338
|
+
os.path.join(*(self.images_folder,) if self.images_folder is not None else (), *(row["image_path"],) if "image_path" in row else (), row[f"image{i}"])
|
|
339
|
+
for i in range(1, self.num_image_cols + 1)]
|
|
340
|
+
|
|
341
|
+
|
|
342
|
+
class StimuliLoader():
|
|
343
|
+
def __init__(self, dataset_file=None, batch_size=1, image_folder=None, id_column=None,
|
|
344
|
+
stimuli_type='lookit', pairwise=False, transform=None):
|
|
345
|
+
self.id_column = id_column
|
|
346
|
+
self.transform = transform
|
|
347
|
+
|
|
348
|
+
if dataset_file is None:
|
|
349
|
+
# If only an image directory is provided
|
|
350
|
+
if image_folder is not None:
|
|
351
|
+
image_files = [
|
|
352
|
+
f for f in os.listdir(image_folder)
|
|
353
|
+
if (
|
|
354
|
+
os.path.isfile(os.path.join(image_folder, f)) and
|
|
355
|
+
not f.startswith("._") and
|
|
356
|
+
f.lower().endswith((".png", ".jpg", ".jpeg"))
|
|
357
|
+
)
|
|
358
|
+
]
|
|
359
|
+
self.manifest = pd.DataFrame({'image1': image_files})
|
|
360
|
+
# assuming that the image path is a unique identifier if a specific ID column is not provided
|
|
361
|
+
self.id_column = "image1" if id_column is None else id_column
|
|
362
|
+
else:
|
|
363
|
+
raise ValueError("Either image folder or dataset file needs to be provided")
|
|
364
|
+
else:
|
|
365
|
+
self.manifest = process_csv(dataset_file)
|
|
366
|
+
|
|
367
|
+
self.stimuli_type = stimuli_type
|
|
368
|
+
self.image_folder = image_folder
|
|
369
|
+
self.batch_size = batch_size
|
|
370
|
+
self.pairwise = pairwise
|
|
371
|
+
|
|
372
|
+
def collator(self, batch):
|
|
373
|
+
return {key: [item for ex in batch for item in ex[key]] for key in batch[0]}
|
|
374
|
+
|
|
375
|
+
def dataloader(self):
|
|
376
|
+
dataset = StimuliDataset(self.manifest, self.image_folder, self.id_column, self.transform)
|
|
377
|
+
return DataLoader(dataset, batch_size=self.batch_size, collate_fn=self.collator)
|
|
@@ -0,0 +1,20 @@
|
|
|
1
|
+
from transformers import pipeline
|
|
2
|
+
from vislearnlabpy.models.vision_model import VisionModel
|
|
3
|
+
from transformers import AutoImageProcessor, AutoModel
|
|
4
|
+
import torch
|
|
5
|
+
|
|
6
|
+
class DinoV3Generator(VisionModel):
|
|
7
|
+
def __init__(self, dataloader=None, device=None, text_prompt="a photo of a "):
|
|
8
|
+
self.pipe = pipeline(
|
|
9
|
+
task="image-feature-extraction",
|
|
10
|
+
model="facebook/dinov3-vits16-pretrain-lvd1689m",
|
|
11
|
+
dtype=torch.bfloat16,
|
|
12
|
+
)
|
|
13
|
+
self.model, self.preprocess = self.pipe.model, self.pipe.feature_extractor
|
|
14
|
+
super().__init__(self.model, self.preprocess, dataloader, device)
|
|
15
|
+
self.name = "dinov3"
|
|
16
|
+
|
|
17
|
+
def image_embeddings(self):
|
|
18
|
+
#
|
|
19
|
+
return
|
|
20
|
+
|
{vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/models/multimodal_model.py
RENAMED
|
@@ -75,42 +75,3 @@ class MultimodalModel(FeatureGenerator):
|
|
|
75
75
|
luce = distractor_similarity / (distractor_similarity + target_similarity)
|
|
76
76
|
return luce
|
|
77
77
|
|
|
78
|
-
def similarities(self, word1, word2, images):
|
|
79
|
-
valid_images = [img for img in images if img is not None]
|
|
80
|
-
similarity_scores = []
|
|
81
|
-
# TODO: this only returns the similarity scores for the first pair of images: need to separate out, indexing is weird
|
|
82
|
-
for image1, image2 in itertools.combinations(valid_images, 2):
|
|
83
|
-
curr_image_embeddings = self.image_embeddings([image1, image2])
|
|
84
|
-
curr_text_embeddings = self.text_embeddings([word1, word2])
|
|
85
|
-
# TODO: need to fix how each row is labeled in lookit_similarities,
|
|
86
|
-
similarity_scores.append({
|
|
87
|
-
'image_similarity': self.similarity(curr_image_embeddings[0], curr_image_embeddings[1]),
|
|
88
|
-
'text_similarity': self.similarity(curr_text_embeddings[0], curr_text_embeddings[1]),
|
|
89
|
-
# finding distractor image to target word similarity
|
|
90
|
-
'multimodal_similarity': self.text_to_images_similarity(curr_image_embeddings, curr_text_embeddings[0], logit_scale=10),
|
|
91
|
-
})
|
|
92
|
-
if similarity_scores == []:
|
|
93
|
-
print(f"skipping {word1} and {word2} since they do not have valid images")
|
|
94
|
-
return [{
|
|
95
|
-
'image_similarity': None,
|
|
96
|
-
'text_similarity': None,
|
|
97
|
-
'multimodal_similarity': None
|
|
98
|
-
}]
|
|
99
|
-
else:
|
|
100
|
-
return similarity_scores
|
|
101
|
-
|
|
102
|
-
# TODO: probably move this to the dataloader row level instead of to a pair of words within a dataloader row
|
|
103
|
-
# TODO: words or texts? what is my parameter
|
|
104
|
-
def embeddings(self, word1, word2, dataloader_row):
|
|
105
|
-
valid_images = [img for img in dataloader_row['images'] if img is not None]
|
|
106
|
-
output_embeddings = []
|
|
107
|
-
for image1, image2 in itertools.combinations(valid_images, 2):
|
|
108
|
-
curr_image_embeddings = self.image_embeddings([image1, image2])
|
|
109
|
-
curr_text_embeddings = self.text_embeddings([word1, word2])
|
|
110
|
-
output_embeddings.append({
|
|
111
|
-
'image_embeddings': curr_image_embeddings,
|
|
112
|
-
'text_embeddings': curr_text_embeddings,
|
|
113
|
-
'multimodal_embeddings': self.multimodal_embeddings(curr_image_embeddings, curr_text_embeddings)
|
|
114
|
-
})
|
|
115
|
-
return output_embeddings
|
|
116
|
-
|
|
@@ -0,0 +1,54 @@
|
|
|
1
|
+
import itertools
|
|
2
|
+
from vislearnlabpy.models.feature_generator import FeatureGenerator
|
|
3
|
+
from vislearnlabpy.embeddings import utils
|
|
4
|
+
from torchvision import transforms
|
|
5
|
+
import torch
|
|
6
|
+
|
|
7
|
+
class VisionModel(FeatureGenerator):
|
|
8
|
+
"""Abstract base class for multimodal models like CLIP and CVCL that extends FeatureGenerator"""
|
|
9
|
+
|
|
10
|
+
def __init__(self, model, preprocess, dataloader=None, device=None):
|
|
11
|
+
super().__init__(model, preprocess, dataloader, device)
|
|
12
|
+
self.image_word_alignment = lambda **x: self.model(**x).logits_per_image.softmax(dim=-1).detach().cpu().numpy()
|
|
13
|
+
|
|
14
|
+
# Load and preprocess images
|
|
15
|
+
def preprocess_image(self, image):
|
|
16
|
+
if isinstance(image, torch.Tensor):
|
|
17
|
+
transform = transforms.ToPILImage()
|
|
18
|
+
image = transform(image)
|
|
19
|
+
return self.preprocess(image).unsqueeze(0).to(self.device)
|
|
20
|
+
|
|
21
|
+
def encode_image(self, image):
|
|
22
|
+
return self.model.encode_image(image)
|
|
23
|
+
|
|
24
|
+
def image_embeddings(self, images, normalize_embeddings=False):
|
|
25
|
+
"""Get image embeddings (batched for speed)"""
|
|
26
|
+
# Handle single image case
|
|
27
|
+
if not isinstance(images, list):
|
|
28
|
+
return self.image_embeddings([images], normalize_embeddings)[0]
|
|
29
|
+
# Preprocess all images → batch
|
|
30
|
+
preprocessed_images = [self.preprocess_image(image) for image in images]
|
|
31
|
+
preprocessed_images = [image.squeeze(0) if image.dim() == 4 else image for image in preprocessed_images]
|
|
32
|
+
# Stack into a single tensor batch (assuming tensors are returned)
|
|
33
|
+
image_batch = torch.stack(preprocessed_images).to(self.device)
|
|
34
|
+
with torch.no_grad():
|
|
35
|
+
embeddings = self.encode_image(image_batch) # model handles batch
|
|
36
|
+
if normalize_embeddings:
|
|
37
|
+
embeddings = utils.normalize_embeddings(embeddings)
|
|
38
|
+
return embeddings
|
|
39
|
+
|
|
40
|
+
# TODO: probably move this to the dataloader row level instead of to a pair of words within a dataloader row
|
|
41
|
+
# TODO: words or texts? what is my parameter
|
|
42
|
+
def embeddings(self, word1, word2, dataloader_row):
|
|
43
|
+
valid_images = [img for img in dataloader_row['images'] if img is not None]
|
|
44
|
+
output_embeddings = []
|
|
45
|
+
for image1, image2 in itertools.combinations(valid_images, 2):
|
|
46
|
+
curr_image_embeddings = self.image_embeddings([image1, image2])
|
|
47
|
+
curr_text_embeddings = self.text_embeddings([word1, word2])
|
|
48
|
+
output_embeddings.append({
|
|
49
|
+
'image_embeddings': curr_image_embeddings,
|
|
50
|
+
'text_embeddings': curr_text_embeddings,
|
|
51
|
+
'multimodal_embeddings': self.multimodal_embeddings(curr_image_embeddings, curr_text_embeddings)
|
|
52
|
+
})
|
|
53
|
+
return output_embeddings
|
|
54
|
+
|
|
@@ -18,5 +18,7 @@ src/vislearnlabpy/embeddings/utils.py
|
|
|
18
18
|
src/vislearnlabpy/extractions/drawingtask_extractor.py
|
|
19
19
|
src/vislearnlabpy/extractions/mongo_extractor.py
|
|
20
20
|
src/vislearnlabpy/models/clip_model.py
|
|
21
|
+
src/vislearnlabpy/models/dinov3_model.py
|
|
21
22
|
src/vislearnlabpy/models/feature_generator.py
|
|
22
|
-
src/vislearnlabpy/models/multimodal_model.py
|
|
23
|
+
src/vislearnlabpy/models/multimodal_model.py
|
|
24
|
+
src/vislearnlabpy/models/vision_model.py
|
|
@@ -1,207 +0,0 @@
|
|
|
1
|
-
import os
|
|
2
|
-
import pandas as pd
|
|
3
|
-
import re
|
|
4
|
-
from PIL import Image
|
|
5
|
-
from torch.utils.data import Dataset, DataLoader
|
|
6
|
-
from vislearnlabpy.embeddings.utils import process_csv
|
|
7
|
-
import random
|
|
8
|
-
from typing import Callable, Any
|
|
9
|
-
import warnings
|
|
10
|
-
from torchvision import transforms
|
|
11
|
-
import numpy as np
|
|
12
|
-
|
|
13
|
-
random.seed(2)
|
|
14
|
-
|
|
15
|
-
class ImageExtractor:
|
|
16
|
-
def __init__(self):
|
|
17
|
-
pass
|
|
18
|
-
|
|
19
|
-
@staticmethod
|
|
20
|
-
def RGBA2RGB(img, background_color=(255, 255, 255)):
|
|
21
|
-
"""Alpha composite an RGBA Image with a specified color.
|
|
22
|
-
Source: http://stackoverflow.com/a/9459208/284318
|
|
23
|
-
|
|
24
|
-
Keyword Arguments:
|
|
25
|
-
img -- PIL RGBA Image object
|
|
26
|
-
background_color -- Tuple r, g, b (default 255, 255, 255)
|
|
27
|
-
"""
|
|
28
|
-
if img.mode in ('RGBA', 'LA'): # Image has transparency
|
|
29
|
-
# Create white background
|
|
30
|
-
white_bg = Image.new('RGB', img.size, background_color)
|
|
31
|
-
# Paste the image onto white background
|
|
32
|
-
# The image itself acts as the alpha mask
|
|
33
|
-
white_bg.paste(img, mask=img.split()[-1] if img.mode == 'RGBA' else None)
|
|
34
|
-
img = white_bg
|
|
35
|
-
elif img.mode == 'P': # Palette mode, might have transparency
|
|
36
|
-
img = img.convert('RGBA')
|
|
37
|
-
white_bg = Image.new('RGB', img.size, background_color)
|
|
38
|
-
white_bg.paste(img, mask=img.split()[-1])
|
|
39
|
-
img = white_bg
|
|
40
|
-
else:
|
|
41
|
-
# Image is already RGB/L, convert to RGB if needed
|
|
42
|
-
img = img.convert('RGB')
|
|
43
|
-
return img
|
|
44
|
-
|
|
45
|
-
@staticmethod
|
|
46
|
-
def crop_to_content(img, apply_content_crop):
|
|
47
|
-
"""Crop image to remove white space around content."""
|
|
48
|
-
if not apply_content_crop:
|
|
49
|
-
return img
|
|
50
|
-
|
|
51
|
-
arr = np.asarray(img)
|
|
52
|
-
rows, cols, channels = np.where(arr < 255)
|
|
53
|
-
|
|
54
|
-
try:
|
|
55
|
-
xlb = min(cols) # left bound
|
|
56
|
-
xub = max(cols) # right bound
|
|
57
|
-
ylb = min(rows) # top bound
|
|
58
|
-
yub = max(rows) # bottom bound
|
|
59
|
-
|
|
60
|
-
# Make it square by using the same bounds
|
|
61
|
-
lb = min([xlb, ylb])
|
|
62
|
-
ub = max([xub, yub])
|
|
63
|
-
|
|
64
|
-
img = img.crop((lb, lb, ub, ub))
|
|
65
|
-
except ValueError:
|
|
66
|
-
print('Blank image - skipping crop')
|
|
67
|
-
return img
|
|
68
|
-
|
|
69
|
-
@staticmethod
|
|
70
|
-
def get_transformations(resize_dim=256, crop_dim=224, apply_content_crop=True, apply_center_crop=False, use_thumbnail=False):
|
|
71
|
-
"""Load image transformations for dataloader.
|
|
72
|
-
|
|
73
|
-
Args:
|
|
74
|
-
resize_dim: Dimension for resizing (default 256)
|
|
75
|
-
crop_dim: Dimension for center crop (default 224)
|
|
76
|
-
apply_content_crop: Whether to crop to content (default True)
|
|
77
|
-
apply_center_crop: Whether to apply center crop (default True)
|
|
78
|
-
use_thumbnail: Whether to use thumbnail resizing instead of regular resize (default False)
|
|
79
|
-
"""
|
|
80
|
-
|
|
81
|
-
def combined_transform(image):
|
|
82
|
-
# Step 1: Apply thumbnail resizing if requested
|
|
83
|
-
if use_thumbnail:
|
|
84
|
-
image.thumbnail((resize_dim, resize_dim), Image.Resampling.LANCZOS)
|
|
85
|
-
|
|
86
|
-
# Step 2: Convert RGBA to RGB
|
|
87
|
-
img_rgb = ImageExtractor.RGBA2RGB(image)
|
|
88
|
-
|
|
89
|
-
# Step 3: Crop to content (remove whitespace) if enabled
|
|
90
|
-
img_cropped = ImageExtractor.crop_to_content(img_rgb, apply_content_crop)
|
|
91
|
-
return img_cropped
|
|
92
|
-
|
|
93
|
-
# Build transformation pipeline
|
|
94
|
-
transform_list = []
|
|
95
|
-
|
|
96
|
-
# Only add regular resize if not using thumbnail
|
|
97
|
-
if not use_thumbnail:
|
|
98
|
-
# resize first
|
|
99
|
-
transform_list.append(transforms.Resize(resize_dim))
|
|
100
|
-
|
|
101
|
-
transform_list.append(transforms.Lambda(combined_transform))
|
|
102
|
-
|
|
103
|
-
if apply_center_crop:
|
|
104
|
-
transform_list.append(transforms.CenterCrop(crop_dim))
|
|
105
|
-
|
|
106
|
-
# Add any additional transforms you might need
|
|
107
|
-
#transform_list.extend([
|
|
108
|
-
#transforms.ToTensor(),
|
|
109
|
-
# transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]) # ImageNet normalization
|
|
110
|
-
#])
|
|
111
|
-
|
|
112
|
-
return transforms.Compose(transform_list)
|
|
113
|
-
|
|
114
|
-
class StimuliDataset(Dataset):
|
|
115
|
-
def __init__(self, manifest, images_folder=None, id_column=None, transform=None):
|
|
116
|
-
self.manifest = manifest
|
|
117
|
-
self.images_folder = images_folder
|
|
118
|
-
self.num_text_cols = sum(1 for c in self.manifest.columns if re.compile("text[0-9]").match(c))
|
|
119
|
-
self.num_image_cols = len([c for c in self.manifest.columns if re.compile("image[0-9]").match(c)])
|
|
120
|
-
self.id_column = id_column
|
|
121
|
-
self.transform = transform
|
|
122
|
-
|
|
123
|
-
def __len__(self):
|
|
124
|
-
return len(self.manifest)
|
|
125
|
-
|
|
126
|
-
def __getitem__(self, idx):
|
|
127
|
-
row = self.manifest.iloc[idx]
|
|
128
|
-
texts = [str(row[f"text{i}"]) for i in range(1, self.num_text_cols + 1)]
|
|
129
|
-
images = []
|
|
130
|
-
image_paths = self._get_image_paths(row)
|
|
131
|
-
full_image_paths = []
|
|
132
|
-
|
|
133
|
-
for image_path in image_paths:
|
|
134
|
-
if not os.path.isabs(image_path):
|
|
135
|
-
image_path = os.path.join(os.getcwd(), image_path)
|
|
136
|
-
if image_path is None or not os.path.exists(image_path):
|
|
137
|
-
full_image_paths.append(None)
|
|
138
|
-
images.append(None)
|
|
139
|
-
else:
|
|
140
|
-
with Image.open(image_path) as img:
|
|
141
|
-
img_copy = img.copy() # Copy the image data to memory
|
|
142
|
-
|
|
143
|
-
if self.transform is not None:
|
|
144
|
-
# Apply the transform pipeline
|
|
145
|
-
img_copy = self.transform(img_copy)
|
|
146
|
-
else:
|
|
147
|
-
# Fallback to just converting to RGB
|
|
148
|
-
img_copy = img_copy.convert('RGB')
|
|
149
|
-
|
|
150
|
-
full_image_paths.append(image_path)
|
|
151
|
-
images.append(img_copy)
|
|
152
|
-
|
|
153
|
-
# just assuming the image path as an item id for now if there is more than one item per row.
|
|
154
|
-
# TODO: allow for the provision of both item ids and row ids
|
|
155
|
-
row_id = [row[self.id_column] if self.id_column is not None else random.randint(0, 1000000)]
|
|
156
|
-
return {"images": images, "text": texts, "item_id": full_image_paths if len(images) > 0 else row_id, "row_id": row_id}
|
|
157
|
-
|
|
158
|
-
def _get_image_paths(self, row):
|
|
159
|
-
# If this is a text only dataset
|
|
160
|
-
if self.images_folder is None and self.num_image_cols == 0:
|
|
161
|
-
return [None] * self.num_text_cols
|
|
162
|
-
# If images are not in the manifest
|
|
163
|
-
elif self.num_image_cols == 0:
|
|
164
|
-
return [os.path.join(self.images_folder, row[f"text{i}"] + ".jpg") for i in range(1, self.num_text_cols + 1)]
|
|
165
|
-
else:
|
|
166
|
-
# allowing for a secondary image path to be provided with different images in the same stimuli sets stored in different subpaths
|
|
167
|
-
return [
|
|
168
|
-
os.path.join(*(self.images_folder,) if self.images_folder is not None else (), *(row["image_path"],) if "image_path" in row else (), row[f"image{i}"])
|
|
169
|
-
for i in range(1, self.num_image_cols + 1)]
|
|
170
|
-
|
|
171
|
-
|
|
172
|
-
class StimuliLoader():
|
|
173
|
-
def __init__(self, dataset_file=None, batch_size=1, image_folder=None, id_column=None,
|
|
174
|
-
stimuli_type='lookit', pairwise=False, transform=None):
|
|
175
|
-
self.id_column = id_column
|
|
176
|
-
self.transform = transform
|
|
177
|
-
|
|
178
|
-
if dataset_file is None:
|
|
179
|
-
# If only an image directory is provided
|
|
180
|
-
if image_folder is not None:
|
|
181
|
-
image_files = [
|
|
182
|
-
f for f in os.listdir(image_folder)
|
|
183
|
-
if (
|
|
184
|
-
os.path.isfile(os.path.join(image_folder, f)) and
|
|
185
|
-
not f.startswith("._") and
|
|
186
|
-
f.lower().endswith((".png", ".jpg", ".jpeg"))
|
|
187
|
-
)
|
|
188
|
-
]
|
|
189
|
-
self.manifest = pd.DataFrame({'image1': image_files})
|
|
190
|
-
# assuming that the image path is a unique identifier if a specific ID column is not provided
|
|
191
|
-
self.id_column = "image1" if id_column is None else id_column
|
|
192
|
-
else:
|
|
193
|
-
raise ValueError("Either image folder or dataset file needs to be provided")
|
|
194
|
-
else:
|
|
195
|
-
self.manifest = process_csv(dataset_file)
|
|
196
|
-
|
|
197
|
-
self.stimuli_type = stimuli_type
|
|
198
|
-
self.image_folder = image_folder
|
|
199
|
-
self.batch_size = batch_size
|
|
200
|
-
self.pairwise = pairwise
|
|
201
|
-
|
|
202
|
-
def collator(self, batch):
|
|
203
|
-
return {key: [item for ex in batch for item in ex[key]] for key in batch[0]}
|
|
204
|
-
|
|
205
|
-
def dataloader(self):
|
|
206
|
-
dataset = StimuliDataset(self.manifest, self.image_folder, self.id_column, self.transform)
|
|
207
|
-
return DataLoader(dataset, batch_size=self.batch_size, collate_fn=self.collator)
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/drawings/svg_render_helpers.py
RENAMED
|
File without changes
|
{vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/embeddings/embedding_store.py
RENAMED
|
File without changes
|
{vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/embeddings/generate_embeddings.py
RENAMED
|
File without changes
|
{vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/embeddings/similarity_generator.py
RENAMED
|
File without changes
|
{vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/embeddings/similarity_utils.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
{vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/extractions/mongo_extractor.py
RENAMED
|
File without changes
|
|
File without changes
|
{vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/models/feature_generator.py
RENAMED
|
File without changes
|
{vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy.egg-info/dependency_links.txt
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|