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.
Files changed (27) hide show
  1. {vislearnlabpy-0.0.2.2/src/vislearnlabpy.egg-info → vislearnlabpy-0.0.2.4}/PKG-INFO +1 -1
  2. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/pyproject.toml +1 -1
  3. vislearnlabpy-0.0.2.4/src/vislearnlabpy/embeddings/stimuli_loader.py +377 -0
  4. vislearnlabpy-0.0.2.4/src/vislearnlabpy/models/dinov3_model.py +20 -0
  5. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/models/multimodal_model.py +0 -39
  6. vislearnlabpy-0.0.2.4/src/vislearnlabpy/models/vision_model.py +54 -0
  7. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4/src/vislearnlabpy.egg-info}/PKG-INFO +1 -1
  8. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy.egg-info/SOURCES.txt +3 -1
  9. vislearnlabpy-0.0.2.2/src/vislearnlabpy/embeddings/stimuli_loader.py +0 -207
  10. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/LICENSE +0 -0
  11. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/README.md +0 -0
  12. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/setup.cfg +0 -0
  13. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/__init__.py +0 -0
  14. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/drawings/drawing.py +0 -0
  15. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/drawings/svg_render_helpers.py +0 -0
  16. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/embeddings/embedding_store.py +0 -0
  17. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/embeddings/generate_embeddings.py +0 -0
  18. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/embeddings/similarity_generator.py +0 -0
  19. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/embeddings/similarity_utils.py +0 -0
  20. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/embeddings/utils.py +0 -0
  21. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/extractions/drawingtask_extractor.py +0 -0
  22. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/extractions/mongo_extractor.py +0 -0
  23. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/models/clip_model.py +0 -0
  24. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy/models/feature_generator.py +0 -0
  25. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy.egg-info/dependency_links.txt +0 -0
  26. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy.egg-info/requires.txt +0 -0
  27. {vislearnlabpy-0.0.2.2 → vislearnlabpy-0.0.2.4}/src/vislearnlabpy.egg-info/top_level.txt +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: vislearnlabpy
3
- Version: 0.0.2.2
3
+ Version: 0.0.2.4
4
4
  Summary: Visual Learning Lab utility files and pipelines
5
5
  Author-email: Tarun Sepuri <tarunsepuri@gmail.com>
6
6
  License-Expression: MIT
@@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
4
4
 
5
5
  [project]
6
6
  name = "vislearnlabpy"
7
- version = "0.0.2.2"
7
+ version = "0.0.2.4"
8
8
  authors = [
9
9
  { name="Tarun Sepuri", email="tarunsepuri@gmail.com" },
10
10
  ]
@@ -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
+
@@ -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
+
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.4
2
2
  Name: vislearnlabpy
3
- Version: 0.0.2.2
3
+ Version: 0.0.2.4
4
4
  Summary: Visual Learning Lab utility files and pipelines
5
5
  Author-email: Tarun Sepuri <tarunsepuri@gmail.com>
6
6
  License-Expression: MIT
@@ -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