Generate image embeddings for images in a directory. Args: image_dir (str): The directory containing images. extractor: The feature extractor for images. model: The model used for generating embeddings. batchsize (int): The batch size for processing images.
(
image_dir: str, extractor, model, batchsize: int = 16
)
| 167 | |
| 168 | |
| 169 | def get_image_embedding( |
| 170 | image_dir: str, extractor, model, batchsize: int = 16 |
| 171 | ) -> dict[str, torch.Tensor]: |
| 172 | """ |
| 173 | Generate image embeddings for images in a directory. |
| 174 | |
| 175 | Args: |
| 176 | image_dir (str): The directory containing images. |
| 177 | extractor: The feature extractor for images. |
| 178 | model: The model used for generating embeddings. |
| 179 | batchsize (int): The batch size for processing images. |
| 180 | |
| 181 | Returns: |
| 182 | dict: A dictionary mapping image filenames to their embeddings. |
| 183 | """ |
| 184 | transform = T.Compose( |
| 185 | [ |
| 186 | T.Resize(int((256 / 224) * extractor.size["height"])), |
| 187 | T.CenterCrop(extractor.size["height"]), |
| 188 | T.ToTensor(), |
| 189 | T.Normalize(mean=extractor.image_mean, std=extractor.image_std), |
| 190 | ] |
| 191 | ) |
| 192 | |
| 193 | inputs = [] |
| 194 | embeddings = [] |
| 195 | images = [i for i in sorted(os.listdir(image_dir)) if is_image_path(i)] |
| 196 | for file in images: |
| 197 | image = Image.open(pjoin(image_dir, file)).convert("RGB") |
| 198 | inputs.append(transform(image)) |
| 199 | if len(inputs) % batchsize == 0 or file == images[-1]: |
| 200 | batch = {"pixel_values": torch.stack(inputs).to(model.device)} |
| 201 | embeddings.extend(model(**batch).last_hidden_state.detach()) |
| 202 | inputs.clear() |
| 203 | return {image: embedding.flatten() for image, embedding in zip(images, embeddings)} |
| 204 | |
| 205 | |
| 206 | def images_cosine_similarity(embeddings: list[torch.Tensor]) -> torch.Tensor: |
no test coverage detected