evalwise 0.1.0__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.
@@ -0,0 +1,382 @@
1
+ """
2
+ Image assertions - deterministic checks for generated images.
3
+
4
+ These use vision models (CLIP, object detection, etc.) which are
5
+ deterministic given the model weights.
6
+ """
7
+
8
+ from __future__ import annotations
9
+
10
+ import time
11
+ from pathlib import Path
12
+ from typing import Any
13
+
14
+ from evalwise.core.context import get_context
15
+ from evalwise.core.result import AssertionResult, Status
16
+
17
+
18
+ def _load_image(image: str | Path | Any) -> Any:
19
+ """Load image from path or return PIL Image."""
20
+ try:
21
+ from PIL import Image
22
+ except ImportError:
23
+ raise ImportError("Pillow required: pip install pillow")
24
+
25
+ if isinstance(image, (str, Path)):
26
+ return Image.open(image)
27
+ return image
28
+
29
+
30
+ class ImageAssert:
31
+ """
32
+ Deterministic assertions for generated images.
33
+
34
+ Example:
35
+ from evalwise.image import ImageAssert
36
+
37
+ ImageAssert.clip_score(image, "a cat sitting on a couch", threshold=0.25)
38
+ ImageAssert.contains_object(image, "cat")
39
+ ImageAssert.nsfw_below(image, threshold=0.1)
40
+ """
41
+
42
+ # ==================== CLIP-based ====================
43
+
44
+ @staticmethod
45
+ def clip_score(image: str | Path, prompt: str, *, threshold: float = 0.25) -> bool:
46
+ """
47
+ Assert image matches text prompt (CLIP similarity).
48
+
49
+ Args:
50
+ image: Path to image or PIL Image
51
+ prompt: Text description
52
+ threshold: Minimum similarity score (0-1, typically 0.2-0.35 is good)
53
+ """
54
+ start = time.perf_counter()
55
+
56
+ try:
57
+ import torch
58
+ import open_clip
59
+ except ImportError:
60
+ raise ImportError("open-clip-torch required: pip install open-clip-torch")
61
+
62
+ # Load model
63
+ model, _, preprocess = open_clip.create_model_and_transforms(
64
+ 'ViT-B-32', pretrained='laion2b_s34b_b79k'
65
+ )
66
+ tokenizer = open_clip.get_tokenizer('ViT-B-32')
67
+
68
+ # Load and preprocess image
69
+ img = _load_image(image)
70
+ img_tensor = preprocess(img).unsqueeze(0)
71
+
72
+ # Encode
73
+ with torch.no_grad():
74
+ image_features = model.encode_image(img_tensor)
75
+ text_features = model.encode_text(tokenizer([prompt]))
76
+
77
+ # Normalize
78
+ image_features /= image_features.norm(dim=-1, keepdim=True)
79
+ text_features /= text_features.norm(dim=-1, keepdim=True)
80
+
81
+ # Cosine similarity
82
+ similarity = (image_features @ text_features.T).item()
83
+
84
+ passed = similarity >= threshold
85
+
86
+ result = AssertionResult(
87
+ name="clip_score",
88
+ status=Status.PASS if passed else Status.FAIL,
89
+ message=f"CLIP score {similarity:.3f} < threshold {threshold}" if not passed else None,
90
+ expected=f">= {threshold}",
91
+ actual=round(similarity, 3),
92
+ duration_ms=(time.perf_counter() - start) * 1000,
93
+ )
94
+ get_context().add_assertion(result)
95
+ return passed
96
+
97
+ # ==================== Metadata ====================
98
+
99
+ @staticmethod
100
+ def resolution_is(image: str | Path, *, width: int, height: int) -> bool:
101
+ """Assert image has exact resolution."""
102
+ start = time.perf_counter()
103
+
104
+ img = _load_image(image)
105
+ actual_width, actual_height = img.size
106
+
107
+ passed = actual_width == width and actual_height == height
108
+
109
+ result = AssertionResult(
110
+ name="resolution_is",
111
+ status=Status.PASS if passed else Status.FAIL,
112
+ message=f"Resolution {actual_width}x{actual_height} != {width}x{height}" if not passed else None,
113
+ expected=f"{width}x{height}",
114
+ actual=f"{actual_width}x{actual_height}",
115
+ duration_ms=(time.perf_counter() - start) * 1000,
116
+ )
117
+ get_context().add_assertion(result)
118
+ return passed
119
+
120
+ @staticmethod
121
+ def resolution_min(image: str | Path, *, width: int, height: int) -> bool:
122
+ """Assert image meets minimum resolution."""
123
+ start = time.perf_counter()
124
+
125
+ img = _load_image(image)
126
+ actual_width, actual_height = img.size
127
+
128
+ passed = actual_width >= width and actual_height >= height
129
+
130
+ result = AssertionResult(
131
+ name="resolution_min",
132
+ status=Status.PASS if passed else Status.FAIL,
133
+ message=f"Resolution {actual_width}x{actual_height} below minimum {width}x{height}" if not passed else None,
134
+ expected=f">= {width}x{height}",
135
+ actual=f"{actual_width}x{actual_height}",
136
+ duration_ms=(time.perf_counter() - start) * 1000,
137
+ )
138
+ get_context().add_assertion(result)
139
+ return passed
140
+
141
+ @staticmethod
142
+ def aspect_ratio(image: str | Path, *, ratio: float, tolerance: float = 0.05) -> bool:
143
+ """Assert image has expected aspect ratio (width/height)."""
144
+ start = time.perf_counter()
145
+
146
+ img = _load_image(image)
147
+ width, height = img.size
148
+ actual_ratio = width / height
149
+
150
+ passed = abs(actual_ratio - ratio) <= tolerance
151
+
152
+ result = AssertionResult(
153
+ name="aspect_ratio",
154
+ status=Status.PASS if passed else Status.FAIL,
155
+ message=f"Aspect ratio {actual_ratio:.2f} != {ratio} (±{tolerance})" if not passed else None,
156
+ expected=f"{ratio} (±{tolerance})",
157
+ actual=round(actual_ratio, 2),
158
+ duration_ms=(time.perf_counter() - start) * 1000,
159
+ )
160
+ get_context().add_assertion(result)
161
+ return passed
162
+
163
+ @staticmethod
164
+ def format_is(image: str | Path, expected_format: str) -> bool:
165
+ """Assert image format (PNG, JPEG, WEBP, etc.)."""
166
+ start = time.perf_counter()
167
+
168
+ img = _load_image(image)
169
+ actual_format = img.format or "unknown"
170
+
171
+ passed = actual_format.upper() == expected_format.upper()
172
+
173
+ result = AssertionResult(
174
+ name="format_is",
175
+ status=Status.PASS if passed else Status.FAIL,
176
+ message=f"Format {actual_format} != {expected_format}" if not passed else None,
177
+ expected=expected_format.upper(),
178
+ actual=actual_format,
179
+ duration_ms=(time.perf_counter() - start) * 1000,
180
+ )
181
+ get_context().add_assertion(result)
182
+ return passed
183
+
184
+ # ==================== Object Detection ====================
185
+
186
+ @staticmethod
187
+ def contains_object(
188
+ image: str | Path,
189
+ object_name: str,
190
+ *,
191
+ confidence: float = 0.5,
192
+ model: str = "yolov8n",
193
+ ) -> bool:
194
+ """
195
+ Assert image contains a specific object.
196
+
197
+ Uses YOLO for detection - deterministic given model weights.
198
+ """
199
+ start = time.perf_counter()
200
+
201
+ try:
202
+ from ultralytics import YOLO
203
+ except ImportError:
204
+ raise ImportError("ultralytics required: pip install ultralytics")
205
+
206
+ # Load model
207
+ detector = YOLO(model)
208
+
209
+ # Run detection
210
+ img = _load_image(image)
211
+ results = detector(img, verbose=False)
212
+
213
+ # Check for object
214
+ found = False
215
+ found_conf = 0.0
216
+
217
+ for r in results:
218
+ for box in r.boxes:
219
+ class_name = detector.names[int(box.cls)]
220
+ conf = float(box.conf)
221
+
222
+ if object_name.lower() in class_name.lower() and conf >= confidence:
223
+ found = True
224
+ found_conf = max(found_conf, conf)
225
+
226
+ result = AssertionResult(
227
+ name="contains_object",
228
+ status=Status.PASS if found else Status.FAIL,
229
+ message=f"Object '{object_name}' not found with confidence >= {confidence}" if not found else None,
230
+ expected=f"{object_name} (>= {confidence})",
231
+ actual=f"found with {found_conf:.2f}" if found else "not found",
232
+ duration_ms=(time.perf_counter() - start) * 1000,
233
+ )
234
+ get_context().add_assertion(result)
235
+ return found
236
+
237
+ @staticmethod
238
+ def object_count(
239
+ image: str | Path,
240
+ object_name: str,
241
+ *,
242
+ min_count: int | None = None,
243
+ max_count: int | None = None,
244
+ exactly: int | None = None,
245
+ confidence: float = 0.5,
246
+ ) -> bool:
247
+ """Assert number of objects of a type in image."""
248
+ start = time.perf_counter()
249
+
250
+ try:
251
+ from ultralytics import YOLO
252
+ except ImportError:
253
+ raise ImportError("ultralytics required: pip install ultralytics")
254
+
255
+ detector = YOLO("yolov8n")
256
+ img = _load_image(image)
257
+ results = detector(img, verbose=False)
258
+
259
+ count = 0
260
+ for r in results:
261
+ for box in r.boxes:
262
+ class_name = detector.names[int(box.cls)]
263
+ conf = float(box.conf)
264
+ if object_name.lower() in class_name.lower() and conf >= confidence:
265
+ count += 1
266
+
267
+ if exactly is not None:
268
+ passed = count == exactly
269
+ expected = f"exactly {exactly}"
270
+ elif min_count is not None and max_count is not None:
271
+ passed = min_count <= count <= max_count
272
+ expected = f"{min_count}-{max_count}"
273
+ elif min_count is not None:
274
+ passed = count >= min_count
275
+ expected = f">= {min_count}"
276
+ elif max_count is not None:
277
+ passed = count <= max_count
278
+ expected = f"<= {max_count}"
279
+ else:
280
+ passed = True
281
+ expected = "any"
282
+
283
+ result = AssertionResult(
284
+ name="object_count",
285
+ status=Status.PASS if passed else Status.FAIL,
286
+ message=f"Found {count} '{object_name}', expected {expected}" if not passed else None,
287
+ expected=expected,
288
+ actual=count,
289
+ duration_ms=(time.perf_counter() - start) * 1000,
290
+ )
291
+ get_context().add_assertion(result)
292
+ return passed
293
+
294
+ # ==================== Safety ====================
295
+
296
+ @staticmethod
297
+ def nsfw_below(image: str | Path, *, threshold: float = 0.5) -> bool:
298
+ """
299
+ Assert NSFW score is below threshold.
300
+
301
+ Uses a safety classifier - deterministic given model.
302
+ """
303
+ start = time.perf_counter()
304
+
305
+ try:
306
+ from transformers import pipeline
307
+ except ImportError:
308
+ raise ImportError("transformers required: pip install transformers")
309
+
310
+ classifier = pipeline(
311
+ "image-classification",
312
+ model="Falconsai/nsfw_image_detection",
313
+ )
314
+
315
+ img = _load_image(image)
316
+ results = classifier(img)
317
+
318
+ nsfw_score = 0.0
319
+ for r in results:
320
+ if r["label"].lower() == "nsfw":
321
+ nsfw_score = r["score"]
322
+ break
323
+
324
+ passed = nsfw_score < threshold
325
+
326
+ result = AssertionResult(
327
+ name="nsfw_below",
328
+ status=Status.PASS if passed else Status.FAIL,
329
+ message=f"NSFW score {nsfw_score:.3f} >= threshold {threshold}" if not passed else None,
330
+ expected=f"< {threshold}",
331
+ actual=round(nsfw_score, 3),
332
+ duration_ms=(time.perf_counter() - start) * 1000,
333
+ )
334
+ get_context().add_assertion(result)
335
+ return passed
336
+
337
+ # ==================== Similarity ====================
338
+
339
+ @staticmethod
340
+ def image_similarity(
341
+ image: str | Path,
342
+ reference: str | Path,
343
+ *,
344
+ threshold: float = 0.8,
345
+ ) -> bool:
346
+ """Assert image is similar to reference image (CLIP embedding similarity)."""
347
+ start = time.perf_counter()
348
+
349
+ try:
350
+ import torch
351
+ import open_clip
352
+ except ImportError:
353
+ raise ImportError("open-clip-torch required: pip install open-clip-torch")
354
+
355
+ model, _, preprocess = open_clip.create_model_and_transforms(
356
+ 'ViT-B-32', pretrained='laion2b_s34b_b79k'
357
+ )
358
+
359
+ img1 = preprocess(_load_image(image)).unsqueeze(0)
360
+ img2 = preprocess(_load_image(reference)).unsqueeze(0)
361
+
362
+ with torch.no_grad():
363
+ feat1 = model.encode_image(img1)
364
+ feat2 = model.encode_image(img2)
365
+
366
+ feat1 /= feat1.norm(dim=-1, keepdim=True)
367
+ feat2 /= feat2.norm(dim=-1, keepdim=True)
368
+
369
+ similarity = (feat1 @ feat2.T).item()
370
+
371
+ passed = similarity >= threshold
372
+
373
+ result = AssertionResult(
374
+ name="image_similarity",
375
+ status=Status.PASS if passed else Status.FAIL,
376
+ message=f"Similarity {similarity:.3f} < threshold {threshold}" if not passed else None,
377
+ expected=f">= {threshold}",
378
+ actual=round(similarity, 3),
379
+ duration_ms=(time.perf_counter() - start) * 1000,
380
+ )
381
+ get_context().add_assertion(result)
382
+ return passed
@@ -0,0 +1,5 @@
1
+ """Text assertion module - deterministic checks for text/LLM outputs."""
2
+
3
+ from evalwise.text.assertions import Assert
4
+
5
+ __all__ = ["Assert"]