fastembed-gpu 0.5.1__py3-none-any.whl

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 (47) hide show
  1. fastembed/__init__.py +20 -0
  2. fastembed/common/__init__.py +3 -0
  3. fastembed/common/model_management.py +301 -0
  4. fastembed/common/onnx_model.py +132 -0
  5. fastembed/common/preprocessor_utils.py +82 -0
  6. fastembed/common/types.py +16 -0
  7. fastembed/common/utils.py +55 -0
  8. fastembed/embedding.py +24 -0
  9. fastembed/image/__init__.py +3 -0
  10. fastembed/image/image_embedding.py +97 -0
  11. fastembed/image/image_embedding_base.py +44 -0
  12. fastembed/image/onnx_embedding.py +211 -0
  13. fastembed/image/onnx_image_model.py +131 -0
  14. fastembed/image/transform/functional.py +150 -0
  15. fastembed/image/transform/operators.py +268 -0
  16. fastembed/late_interaction/__init__.py +5 -0
  17. fastembed/late_interaction/colbert.py +256 -0
  18. fastembed/late_interaction/jina_colbert.py +62 -0
  19. fastembed/late_interaction/late_interaction_embedding_base.py +62 -0
  20. fastembed/late_interaction/late_interaction_text_embedding.py +114 -0
  21. fastembed/parallel_processor.py +252 -0
  22. fastembed/rerank/cross_encoder/__init__.py +3 -0
  23. fastembed/rerank/cross_encoder/onnx_text_cross_encoder.py +224 -0
  24. fastembed/rerank/cross_encoder/onnx_text_model.py +150 -0
  25. fastembed/rerank/cross_encoder/text_cross_encoder.py +120 -0
  26. fastembed/rerank/cross_encoder/text_cross_encoder_base.py +58 -0
  27. fastembed/sparse/__init__.py +4 -0
  28. fastembed/sparse/bm25.py +347 -0
  29. fastembed/sparse/bm42.py +340 -0
  30. fastembed/sparse/sparse_embedding_base.py +83 -0
  31. fastembed/sparse/sparse_text_embedding.py +121 -0
  32. fastembed/sparse/splade_pp.py +180 -0
  33. fastembed/sparse/utils/tokenizer.py +120 -0
  34. fastembed/text/__init__.py +3 -0
  35. fastembed/text/clip_embedding.py +54 -0
  36. fastembed/text/e5_onnx_embedding.py +72 -0
  37. fastembed/text/onnx_embedding.py +333 -0
  38. fastembed/text/onnx_text_model.py +145 -0
  39. fastembed/text/pooled_embedding.py +92 -0
  40. fastembed/text/pooled_normalized_embedding.py +125 -0
  41. fastembed/text/text_embedding.py +107 -0
  42. fastembed/text/text_embedding_base.py +62 -0
  43. fastembed_gpu-0.5.1.dist-info/LICENSE +201 -0
  44. fastembed_gpu-0.5.1.dist-info/METADATA +262 -0
  45. fastembed_gpu-0.5.1.dist-info/NOTICE +14 -0
  46. fastembed_gpu-0.5.1.dist-info/RECORD +47 -0
  47. fastembed_gpu-0.5.1.dist-info/WHEEL +4 -0
@@ -0,0 +1,97 @@
1
+ from typing import Any, Iterable, Optional, Sequence, Type
2
+
3
+ import numpy as np
4
+
5
+ from fastembed.common import ImageInput, OnnxProvider
6
+ from fastembed.image.image_embedding_base import ImageEmbeddingBase
7
+ from fastembed.image.onnx_embedding import OnnxImageEmbedding
8
+
9
+
10
+ class ImageEmbedding(ImageEmbeddingBase):
11
+ EMBEDDINGS_REGISTRY: list[Type[ImageEmbeddingBase]] = [OnnxImageEmbedding]
12
+
13
+ @classmethod
14
+ def list_supported_models(cls) -> list[dict[str, Any]]:
15
+ """
16
+ Lists the supported models.
17
+
18
+ Returns:
19
+ list[dict[str, Any]]: A list of dictionaries containing the model information.
20
+
21
+ Example:
22
+ ```
23
+ [
24
+ {
25
+ "model": "Qdrant/clip-ViT-B-32-vision",
26
+ "dim": 512,
27
+ "description": "CLIP vision encoder based on ViT-B/32",
28
+ "license": "mit",
29
+ "size_in_GB": 0.33,
30
+ "sources": {
31
+ "hf": "Qdrant/clip-ViT-B-32-vision",
32
+ },
33
+ "model_file": "model.onnx",
34
+ }
35
+ ]
36
+ ```
37
+ """
38
+ result = []
39
+ for embedding in cls.EMBEDDINGS_REGISTRY:
40
+ result.extend(embedding.list_supported_models())
41
+ return result
42
+
43
+ def __init__(
44
+ self,
45
+ model_name: str,
46
+ cache_dir: Optional[str] = None,
47
+ threads: Optional[int] = None,
48
+ providers: Optional[Sequence[OnnxProvider]] = None,
49
+ cuda: bool = False,
50
+ device_ids: Optional[list[int]] = None,
51
+ lazy_load: bool = False,
52
+ **kwargs,
53
+ ):
54
+ super().__init__(model_name, cache_dir, threads, **kwargs)
55
+ for EMBEDDING_MODEL_TYPE in self.EMBEDDINGS_REGISTRY:
56
+ supported_models = EMBEDDING_MODEL_TYPE.list_supported_models()
57
+ if any(model_name.lower() == model["model"].lower() for model in supported_models):
58
+ self.model = EMBEDDING_MODEL_TYPE(
59
+ model_name,
60
+ cache_dir,
61
+ threads=threads,
62
+ providers=providers,
63
+ cuda=cuda,
64
+ device_ids=device_ids,
65
+ lazy_load=lazy_load,
66
+ **kwargs,
67
+ )
68
+ return
69
+
70
+ raise ValueError(
71
+ f"Model {model_name} is not supported in ImageEmbedding."
72
+ "Please check the supported models using `ImageEmbedding.list_supported_models()`"
73
+ )
74
+
75
+ def embed(
76
+ self,
77
+ images: ImageInput,
78
+ batch_size: int = 16,
79
+ parallel: Optional[int] = None,
80
+ **kwargs,
81
+ ) -> Iterable[np.ndarray]:
82
+ """
83
+ Encode a list of documents into list of embeddings.
84
+ We use mean pooling with attention so that the model can handle variable-length inputs.
85
+
86
+ Args:
87
+ images: Iterator of image paths or single image path to embed
88
+ batch_size: Batch size for encoding -- higher values will use more memory, but be faster
89
+ parallel:
90
+ If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
91
+ If 0, use all available cores.
92
+ If None, don't use data-parallel processing, use default onnxruntime threading instead.
93
+
94
+ Returns:
95
+ List of embeddings, one per document
96
+ """
97
+ yield from self.model.embed(images, batch_size, parallel, **kwargs)
@@ -0,0 +1,44 @@
1
+ from typing import Iterable, Optional
2
+
3
+ import numpy as np
4
+
5
+ from fastembed.common.model_management import ModelManagement
6
+ from fastembed.common.types import ImageInput
7
+
8
+
9
+ class ImageEmbeddingBase(ModelManagement):
10
+ def __init__(
11
+ self,
12
+ model_name: str,
13
+ cache_dir: Optional[str] = None,
14
+ threads: Optional[int] = None,
15
+ **kwargs,
16
+ ):
17
+ self.model_name = model_name
18
+ self.cache_dir = cache_dir
19
+ self.threads = threads
20
+ self._local_files_only = kwargs.pop("local_files_only", False)
21
+
22
+ def embed(
23
+ self,
24
+ images: ImageInput,
25
+ batch_size: int = 16,
26
+ parallel: Optional[int] = None,
27
+ **kwargs,
28
+ ) -> Iterable[np.ndarray]:
29
+ """
30
+ Embeds a list of images into a list of embeddings.
31
+
32
+ Args:
33
+ images: The list of image paths to preprocess and embed.
34
+ batch_size: Batch size for encoding
35
+ parallel:
36
+ If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
37
+ If 0, use all available cores.
38
+ If None, don't use data-parallel processing, use default onnxruntime threading instead.
39
+ **kwargs: Additional keyword argument to pass to the embed method.
40
+
41
+ Yields:
42
+ Iterable[np.ndarray]: The embeddings.
43
+ """
44
+ raise NotImplementedError()
@@ -0,0 +1,211 @@
1
+ from typing import Any, Iterable, Optional, Sequence, Type
2
+
3
+ import numpy as np
4
+
5
+ from fastembed.common import ImageInput, OnnxProvider
6
+ from fastembed.common.onnx_model import OnnxOutputContext
7
+ from fastembed.common.utils import define_cache_dir, normalize
8
+ from fastembed.image.image_embedding_base import ImageEmbeddingBase
9
+ from fastembed.image.onnx_image_model import ImageEmbeddingWorker, OnnxImageModel
10
+
11
+ supported_onnx_models = [
12
+ {
13
+ "model": "Qdrant/clip-ViT-B-32-vision",
14
+ "dim": 512,
15
+ "description": "Image embeddings, Multimodal (text&image), 2021 year",
16
+ "license": "mit",
17
+ "size_in_GB": 0.34,
18
+ "sources": {
19
+ "hf": "Qdrant/clip-ViT-B-32-vision",
20
+ },
21
+ "model_file": "model.onnx",
22
+ },
23
+ {
24
+ "model": "Qdrant/resnet50-onnx",
25
+ "dim": 2048,
26
+ "description": "Image embeddings, Unimodal (image), 2016 year",
27
+ "license": "apache-2.0",
28
+ "size_in_GB": 0.1,
29
+ "sources": {
30
+ "hf": "Qdrant/resnet50-onnx",
31
+ },
32
+ "model_file": "model.onnx",
33
+ },
34
+ {
35
+ "model": "Qdrant/Unicom-ViT-B-16",
36
+ "dim": 768,
37
+ "description": "Image embeddings (more detailed than Unicom-ViT-B-32), Multimodal (text&image), 2023 year",
38
+ "license": "apache-2.0",
39
+ "size_in_GB": 0.82,
40
+ "sources": {
41
+ "hf": "Qdrant/Unicom-ViT-B-16",
42
+ },
43
+ "model_file": "model.onnx",
44
+ },
45
+ {
46
+ "model": "Qdrant/Unicom-ViT-B-32",
47
+ "dim": 512,
48
+ "description": "Image embeddings, Multimodal (text&image), 2023 year",
49
+ "license": "apache-2.0",
50
+ "size_in_GB": 0.48,
51
+ "sources": {
52
+ "hf": "Qdrant/Unicom-ViT-B-32",
53
+ },
54
+ "model_file": "model.onnx",
55
+ },
56
+ {
57
+ "model": "jinaai/jina-clip-v1",
58
+ "dim": 768,
59
+ "description": "Image embeddings, Multimodal (text&image), 2024 year",
60
+ "license": "apache-2.0",
61
+ "size_in_GB": 0.34,
62
+ "sources": {
63
+ "hf": "jinaai/jina-clip-v1",
64
+ },
65
+ "model_file": "onnx/vision_model.onnx",
66
+ },
67
+ ]
68
+
69
+
70
+ class OnnxImageEmbedding(ImageEmbeddingBase, OnnxImageModel[np.ndarray]):
71
+ def __init__(
72
+ self,
73
+ model_name: str,
74
+ cache_dir: Optional[str] = None,
75
+ threads: Optional[int] = None,
76
+ providers: Optional[Sequence[OnnxProvider]] = None,
77
+ cuda: bool = False,
78
+ device_ids: Optional[list[int]] = None,
79
+ lazy_load: bool = False,
80
+ device_id: Optional[int] = None,
81
+ **kwargs,
82
+ ):
83
+ """
84
+ Args:
85
+ model_name (str): The name of the model to use.
86
+ cache_dir (str, optional): The path to the cache directory.
87
+ Can be set using the `FASTEMBED_CACHE_PATH` env variable.
88
+ Defaults to `fastembed_cache` in the system's temp directory.
89
+ threads (int, optional): The number of threads single onnxruntime session can use. Defaults to None.
90
+ providers (Optional[Sequence[OnnxProvider]], optional): The list of onnxruntime providers to use.
91
+ Mutually exclusive with the `cuda` and `device_ids` arguments. Defaults to None.
92
+ cuda (bool, optional): Whether to use cuda for inference. Mutually exclusive with `providers`
93
+ Defaults to False.
94
+ device_ids (Optional[list[int]], optional): The list of device ids to use for data parallel processing in
95
+ workers. Should be used with `cuda=True`, mutually exclusive with `providers`. Defaults to None.
96
+ lazy_load (bool, optional): Whether to load the model during class initialization or on demand.
97
+ Should be set to True when using multiple-gpu and parallel encoding. Defaults to False.
98
+ device_id (Optional[int], optional): The device id to use for loading the model in the worker process.
99
+
100
+ Raises:
101
+ ValueError: If the model_name is not in the format <org>/<model> e.g. BAAI/bge-base-en.
102
+ """
103
+
104
+ super().__init__(model_name, cache_dir, threads, **kwargs)
105
+ self.providers = providers
106
+ self.lazy_load = lazy_load
107
+
108
+ # List of device ids, that can be used for data parallel processing in workers
109
+ self.device_ids = device_ids
110
+ self.cuda = cuda
111
+
112
+ # This device_id will be used if we need to load model in current process
113
+ if device_id is not None:
114
+ self.device_id = device_id
115
+ elif self.device_ids is not None:
116
+ self.device_id = self.device_ids[0]
117
+ else:
118
+ self.device_id = None
119
+
120
+ self.model_description = self._get_model_description(model_name)
121
+ self.cache_dir = define_cache_dir(cache_dir)
122
+ self._model_dir = self.download_model(
123
+ self.model_description, self.cache_dir, local_files_only=self._local_files_only
124
+ )
125
+
126
+ if not self.lazy_load:
127
+ self.load_onnx_model()
128
+
129
+ def load_onnx_model(self) -> None:
130
+ """
131
+ Load the onnx model.
132
+ """
133
+ self._load_onnx_model(
134
+ model_dir=self._model_dir,
135
+ model_file=self.model_description["model_file"],
136
+ threads=self.threads,
137
+ providers=self.providers,
138
+ cuda=self.cuda,
139
+ device_id=self.device_id,
140
+ )
141
+
142
+ @classmethod
143
+ def list_supported_models(cls) -> list[dict[str, Any]]:
144
+ """
145
+ Lists the supported models.
146
+
147
+ Returns:
148
+ list[Dict[str, Any]]: A list of dictionaries containing the model information.
149
+ """
150
+ return supported_onnx_models
151
+
152
+ def embed(
153
+ self,
154
+ images: ImageInput,
155
+ batch_size: int = 16,
156
+ parallel: Optional[int] = None,
157
+ **kwargs,
158
+ ) -> Iterable[np.ndarray]:
159
+ """
160
+ Encode a list of images into list of embeddings.
161
+ We use mean pooling with attention so that the model can handle variable-length inputs.
162
+
163
+ Args:
164
+ images: Iterator of image paths or single image path to embed
165
+ batch_size: Batch size for encoding -- higher values will use more memory, but be faster
166
+ parallel:
167
+ If > 1, data-parallel encoding will be used, recommended for offline encoding of large datasets.
168
+ If 0, use all available cores.
169
+ If None, don't use data-parallel processing, use default onnxruntime threading instead.
170
+
171
+ Returns:
172
+ List of embeddings, one per document
173
+ """
174
+
175
+ yield from self._embed_images(
176
+ model_name=self.model_name,
177
+ cache_dir=str(self.cache_dir),
178
+ images=images,
179
+ batch_size=batch_size,
180
+ parallel=parallel,
181
+ providers=self.providers,
182
+ cuda=self.cuda,
183
+ device_ids=self.device_ids,
184
+ **kwargs,
185
+ )
186
+
187
+ @classmethod
188
+ def _get_worker_class(cls) -> Type["ImageEmbeddingWorker"]:
189
+ return OnnxImageEmbeddingWorker
190
+
191
+ def _preprocess_onnx_input(
192
+ self, onnx_input: dict[str, np.ndarray], **kwargs
193
+ ) -> dict[str, np.ndarray]:
194
+ """
195
+ Preprocess the onnx input.
196
+ """
197
+
198
+ return onnx_input
199
+
200
+ def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[np.ndarray]:
201
+ return normalize(output.model_output).astype(np.float32)
202
+
203
+
204
+ class OnnxImageEmbeddingWorker(ImageEmbeddingWorker):
205
+ def init_embedding(self, model_name: str, cache_dir: str, **kwargs) -> OnnxImageEmbedding:
206
+ return OnnxImageEmbedding(
207
+ model_name=model_name,
208
+ cache_dir=cache_dir,
209
+ threads=1,
210
+ **kwargs,
211
+ )
@@ -0,0 +1,131 @@
1
+ import contextlib
2
+ import os
3
+ from multiprocessing import get_all_start_methods
4
+ from pathlib import Path
5
+ from typing import Any, Iterable, Optional, Sequence, Type
6
+
7
+ import numpy as np
8
+ from PIL import Image
9
+
10
+ from fastembed.common import ImageInput, OnnxProvider
11
+ from fastembed.common.onnx_model import EmbeddingWorker, OnnxModel, OnnxOutputContext, T
12
+ from fastembed.common.preprocessor_utils import load_preprocessor
13
+ from fastembed.common.utils import iter_batch
14
+ from fastembed.parallel_processor import ParallelWorkerPool
15
+
16
+ # Holds type of the embedding result
17
+
18
+
19
+ class OnnxImageModel(OnnxModel[T]):
20
+ @classmethod
21
+ def _get_worker_class(cls) -> Type["ImageEmbeddingWorker"]:
22
+ raise NotImplementedError("Subclasses must implement this method")
23
+
24
+ def _post_process_onnx_output(self, output: OnnxOutputContext) -> Iterable[T]:
25
+ raise NotImplementedError("Subclasses must implement this method")
26
+
27
+ def __init__(self) -> None:
28
+ super().__init__()
29
+ self.processor = None
30
+
31
+ def _preprocess_onnx_input(
32
+ self, onnx_input: dict[str, np.ndarray], **kwargs
33
+ ) -> dict[str, np.ndarray]:
34
+ """
35
+ Preprocess the onnx input.
36
+ """
37
+ return onnx_input
38
+
39
+ def _load_onnx_model(
40
+ self,
41
+ model_dir: Path,
42
+ model_file: str,
43
+ threads: Optional[int],
44
+ providers: Optional[Sequence[OnnxProvider]] = None,
45
+ cuda: bool = False,
46
+ device_id: Optional[int] = None,
47
+ ) -> None:
48
+ super()._load_onnx_model(
49
+ model_dir=model_dir,
50
+ model_file=model_file,
51
+ threads=threads,
52
+ providers=providers,
53
+ cuda=cuda,
54
+ device_id=device_id,
55
+ )
56
+ self.processor = load_preprocessor(model_dir=model_dir)
57
+
58
+ def load_onnx_model(self) -> None:
59
+ raise NotImplementedError("Subclasses must implement this method")
60
+
61
+ def _build_onnx_input(self, encoded: np.ndarray) -> dict[str, np.ndarray]:
62
+ return {node.name: encoded for node in self.model.get_inputs()}
63
+
64
+ def onnx_embed(self, images: list[ImageInput], **kwargs) -> OnnxOutputContext:
65
+ with contextlib.ExitStack():
66
+ image_files = [
67
+ Image.open(image) if not isinstance(image, Image.Image) else image
68
+ for image in images
69
+ ]
70
+ encoded = self.processor(image_files)
71
+ onnx_input = self._build_onnx_input(encoded)
72
+ onnx_input = self._preprocess_onnx_input(onnx_input)
73
+ model_output = self.model.run(None, onnx_input)
74
+ embeddings = model_output[0].reshape(len(images), -1)
75
+ return OnnxOutputContext(model_output=embeddings)
76
+
77
+ def _embed_images(
78
+ self,
79
+ model_name: str,
80
+ cache_dir: str,
81
+ images: ImageInput,
82
+ batch_size: int = 256,
83
+ parallel: Optional[int] = None,
84
+ providers: Optional[Sequence[OnnxProvider]] = None,
85
+ cuda: bool = False,
86
+ device_ids: Optional[list[int]] = None,
87
+ **kwargs,
88
+ ) -> Iterable[T]:
89
+ is_small = False
90
+
91
+ if isinstance(images, (str, Path, Image.Image)):
92
+ images = [images]
93
+ is_small = True
94
+
95
+ if isinstance(images, list) and len(images) < batch_size:
96
+ is_small = True
97
+
98
+ if parallel is None or is_small:
99
+ if not hasattr(self, "model") or self.model is None:
100
+ self.load_onnx_model()
101
+
102
+ for batch in iter_batch(images, batch_size):
103
+ yield from self._post_process_onnx_output(self.onnx_embed(batch))
104
+ else:
105
+ if parallel == 0:
106
+ parallel = os.cpu_count()
107
+
108
+ start_method = "forkserver" if "forkserver" in get_all_start_methods() else "spawn"
109
+ params = {
110
+ "model_name": model_name,
111
+ "cache_dir": cache_dir,
112
+ "providers": providers,
113
+ **kwargs,
114
+ }
115
+
116
+ pool = ParallelWorkerPool(
117
+ num_workers=parallel or 1,
118
+ worker=self._get_worker_class(),
119
+ cuda=cuda,
120
+ device_ids=device_ids,
121
+ start_method=start_method,
122
+ )
123
+ for batch in pool.ordered_map(iter_batch(images, batch_size), **params):
124
+ yield from self._post_process_onnx_output(batch)
125
+
126
+
127
+ class ImageEmbeddingWorker(EmbeddingWorker):
128
+ def process(self, items: Iterable[tuple[int, Any]]) -> Iterable[tuple[int, Any]]:
129
+ for idx, batch in items:
130
+ embeddings = self.model.onnx_embed(batch)
131
+ yield idx, embeddings
@@ -0,0 +1,150 @@
1
+ from typing import Sized, Union
2
+
3
+ import numpy as np
4
+ from PIL import Image
5
+
6
+
7
+ def convert_to_rgb(image: Image.Image) -> Image.Image:
8
+ if image.mode == "RGB":
9
+ return image
10
+
11
+ image = image.convert("RGB")
12
+ return image
13
+
14
+
15
+ def center_crop(
16
+ image: Union[Image.Image, np.ndarray],
17
+ size: tuple[int, int],
18
+ ) -> np.ndarray:
19
+ if isinstance(image, np.ndarray):
20
+ _, orig_height, orig_width = image.shape
21
+ else:
22
+ orig_height, orig_width = image.height, image.width
23
+ # (H, W, C) -> (C, H, W)
24
+ image = np.array(image).transpose((2, 0, 1))
25
+
26
+ crop_height, crop_width = size
27
+
28
+ # left upper corner (0, 0)
29
+ top = (orig_height - crop_height) // 2
30
+ bottom = top + crop_height
31
+ left = (orig_width - crop_width) // 2
32
+ right = left + crop_width
33
+
34
+ # Check if cropped area is within image boundaries
35
+ if top >= 0 and bottom <= orig_height and left >= 0 and right <= orig_width:
36
+ image = image[..., top:bottom, left:right]
37
+ return image
38
+
39
+ # Padding with zeros
40
+ new_height = max(crop_height, orig_height)
41
+ new_width = max(crop_width, orig_width)
42
+ new_shape = image.shape[:-2] + (new_height, new_width)
43
+ new_image = np.zeros_like(image, shape=new_shape)
44
+
45
+ top_pad = (new_height - orig_height) // 2
46
+ bottom_pad = top_pad + orig_height
47
+ left_pad = (new_width - orig_width) // 2
48
+ right_pad = left_pad + orig_width
49
+ new_image[..., top_pad:bottom_pad, left_pad:right_pad] = image
50
+
51
+ top += top_pad
52
+ bottom += top_pad
53
+ left += left_pad
54
+ right += left_pad
55
+
56
+ new_image = new_image[
57
+ ..., max(0, top) : min(new_height, bottom), max(0, left) : min(new_width, right)
58
+ ]
59
+
60
+ return new_image
61
+
62
+
63
+ def normalize(
64
+ image: np.ndarray,
65
+ mean: Union[float, np.ndarray],
66
+ std: Union[float, np.ndarray],
67
+ ) -> np.ndarray:
68
+ if not isinstance(image, np.ndarray):
69
+ raise ValueError("image must be a numpy array")
70
+
71
+ num_channels = image.shape[1] if len(image.shape) == 4 else image.shape[0]
72
+
73
+ if not np.issubdtype(image.dtype, np.floating):
74
+ image = image.astype(np.float32)
75
+
76
+ if isinstance(mean, Sized):
77
+ if len(mean) != num_channels:
78
+ raise ValueError(
79
+ f"mean must have {num_channels} elements if it is an iterable, got {len(mean)}"
80
+ )
81
+ else:
82
+ mean = [mean] * num_channels
83
+ mean = np.array(mean, dtype=image.dtype)
84
+
85
+ if isinstance(std, Sized):
86
+ if len(std) != num_channels:
87
+ raise ValueError(
88
+ f"std must have {num_channels} elements if it is an iterable, got {len(std)}"
89
+ )
90
+ else:
91
+ std = [std] * num_channels
92
+ std = np.array(std, dtype=image.dtype)
93
+
94
+ image = ((image.T - mean) / std).T
95
+ return image
96
+
97
+
98
+ def resize(
99
+ image: Image.Image,
100
+ size: Union[int, tuple[int, int]],
101
+ resample: Union[int, Image.Resampling] = Image.Resampling.BILINEAR,
102
+ ) -> Image.Image:
103
+ if isinstance(size, tuple):
104
+ return image.resize(size, resample)
105
+
106
+ height, width = image.height, image.width
107
+ short, long = (width, height) if width <= height else (height, width)
108
+
109
+ new_short, new_long = size, int(size * long / short)
110
+ if width <= height:
111
+ new_size = (new_short, new_long)
112
+ else:
113
+ new_size = (new_long, new_short)
114
+ return image.resize(new_size, resample)
115
+
116
+
117
+ def rescale(image: np.ndarray, scale: float, dtype=np.float32) -> np.ndarray:
118
+ return (image * scale).astype(dtype)
119
+
120
+
121
+ def pil2ndarray(image: Union[Image.Image, np.ndarray]):
122
+ if isinstance(image, Image.Image):
123
+ return np.asarray(image).transpose((2, 0, 1))
124
+ return image
125
+
126
+
127
+ def pad2square(
128
+ image: Image.Image,
129
+ size: int,
130
+ fill_color: Union[str, int, tuple[int, ...]] = 0,
131
+ ) -> Image.Image:
132
+ height, width = image.height, image.width
133
+
134
+ left, right = 0, width
135
+ top, bottom = 0, height
136
+
137
+ crop_required = False
138
+ if width > size:
139
+ left = (width - size) // 2
140
+ right = left + size
141
+ crop_required = True
142
+
143
+ if height > size:
144
+ top = (height - size) // 2
145
+ bottom = top + size
146
+ crop_required = True
147
+
148
+ new_image = Image.new(mode="RGB", size=(size, size), color=fill_color)
149
+ new_image.paste(image.crop((left, top, right, bottom)) if crop_required else image)
150
+ return new_image