meta-sam-parser 0.0.2__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,533 @@
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates. All Rights Reserved.
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+ import re
7
+ from dataclasses import dataclass
8
+ from typing import cast
9
+
10
+ from ._errors import InvalidSegmentationMaskError
11
+ from ._mask_codec import decode_mask_to_raster
12
+ from ._types import (
13
+ FrameReference,
14
+ ImageSegmentationResult,
15
+ ImageSegmentationSnapshot,
16
+ ParserFinish,
17
+ ResponseFormat,
18
+ ResponseStreamOutcome,
19
+ SegmentationBoxRecord,
20
+ SegmentationDiagnostic,
21
+ SegmentationMask,
22
+ SegmentationMaskBounds,
23
+ SegmentationMaskIdentity,
24
+ SegmentationMaskRecord,
25
+ SegmentationMedia,
26
+ SegmentationPointRecord,
27
+ SegmentationRecord,
28
+ SegmentationResult,
29
+ SegmentationSnapshot,
30
+ SegmentationTextRecord,
31
+ VideoSegmentationResult,
32
+ VideoSegmentationSnapshot,
33
+ )
34
+
35
+ _MAXIMUM_SAFE_INTEGER = (1 << 53) - 1
36
+ _JS_WHITESPACE = (
37
+ "\t\n\v\f\r \u00a0\u1680\u2000\u2001\u2002\u2003\u2004\u2005\u2006"
38
+ "\u2007\u2008\u2009\u200a\u2028\u2029\u202f\u205f\u3000\ufeff"
39
+ )
40
+ _JS_WHITESPACE_CLASS = (
41
+ r"\t\n\v\f\r \u00a0\u1680\u2000-\u200a\u2028\u2029\u202f\u205f\u3000\ufeff"
42
+ )
43
+ _JS_WHITESPACE_PATTERN = rf"[{_JS_WHITESPACE_CLASS}]"
44
+ _JS_NON_WHITESPACE_PATTERN = rf"[^{_JS_WHITESPACE_CLASS}]"
45
+ _JS_DOT_PATTERN = r"[^\n\r\u2028\u2029]"
46
+ _RECORD_PATTERN = re.compile(
47
+ rf"^(?:\[frame=([0-9]+)\]{_JS_WHITESPACE_PATTERN}+)?"
48
+ rf"object=([A-Za-z0-9._-]+){_JS_WHITESPACE_PATTERN}+"
49
+ rf"(point|box|mask)=({_JS_DOT_PATTERN}+)$"
50
+ )
51
+ _COMPACT_API_HEADER_PATTERN = re.compile(rf"^<([0-9]+)f>({_JS_DOT_PATTERN}*)$")
52
+ _COMPACT_API_OBJECT_PATTERN = re.compile(
53
+ r"^(?:,)?([0-9]+)"
54
+ r"<\|box;x1=(-?[0-9]+);y1=(-?[0-9]+);x2=(-?[0-9]+);"
55
+ r"y2=(-?[0-9]+);w=([0-9]+);h=([0-9]+)\|>"
56
+ r"<\|mask;x=0;y=0;data=([0-9]+),([0-9]+),([!~][^|]+)\|>"
57
+ )
58
+ _POINT_PATTERN = re.compile(r"^\(([^,]+),([^,]+)\)$")
59
+ _BOX_PATTERN = re.compile(r"^\(([^,]+),([^,]+),([^,]+),([^,]+)\)$")
60
+ _MASK_PATTERN = re.compile(
61
+ rf"^([^;{_JS_WHITESPACE_CLASS}]+);"
62
+ rf"size=([0-9]+)x([0-9]+);data=({_JS_NON_WHITESPACE_PATTERN}+)$"
63
+ )
64
+ _DECIMAL_NUMBER_PATTERN = re.compile(
65
+ r"^[+-]?(?:(?:[0-9]+(?:\.[0-9]*)?)|(?:\.[0-9]+))"
66
+ r"(?:[eE][+-]?[0-9]+)?$"
67
+ )
68
+ _HEXADECIMAL_NUMBER_PATTERN = re.compile(r"^0[xX][0-9A-Fa-f]+$")
69
+ _BINARY_NUMBER_PATTERN = re.compile(r"^0[bB][01]+$")
70
+ _OCTAL_NUMBER_PATTERN = re.compile(r"^0[oO][0-7]+$")
71
+
72
+
73
+ def _trim(value: str) -> str:
74
+ return value.strip(_JS_WHITESPACE)
75
+
76
+
77
+ def _javascript_number(value: str) -> float | None:
78
+ normalized = _trim(value)
79
+ if not normalized:
80
+ return None
81
+ try:
82
+ if _HEXADECIMAL_NUMBER_PATTERN.fullmatch(normalized):
83
+ parsed = float(int(normalized[2:], 16))
84
+ elif _BINARY_NUMBER_PATTERN.fullmatch(normalized):
85
+ parsed = float(int(normalized[2:], 2))
86
+ elif _OCTAL_NUMBER_PATTERN.fullmatch(normalized):
87
+ parsed = float(int(normalized[2:], 8))
88
+ elif _DECIMAL_NUMBER_PATTERN.fullmatch(normalized):
89
+ parsed = float(normalized)
90
+ else:
91
+ return None
92
+ except (OverflowError, ValueError):
93
+ return None
94
+ return parsed if math.isfinite(parsed) else None
95
+
96
+
97
+ def _integer_text(value: str) -> int | None:
98
+ parsed = _javascript_number(value)
99
+ if parsed is None or not parsed.is_integer() or abs(parsed) > _MAXIMUM_SAFE_INTEGER:
100
+ return None
101
+ return int(parsed)
102
+
103
+
104
+ class _SegmentationParser:
105
+ __slots__ = (
106
+ "_buffer_parts",
107
+ "_diagnostics",
108
+ "_line",
109
+ "_media",
110
+ "_raw_output_parts",
111
+ "_records",
112
+ "_revision",
113
+ "_revisions",
114
+ )
115
+
116
+ def __init__(self, media: SegmentationMedia) -> None:
117
+ self._media = media
118
+ self._records: list[SegmentationRecord] = []
119
+ self._diagnostics: list[SegmentationDiagnostic] = []
120
+ self._revisions: dict[SegmentationMaskIdentity, int] = {}
121
+ self._raw_output_parts: list[str] = []
122
+ self._buffer_parts: list[str] = []
123
+ self._line = 0
124
+ self._revision = 0
125
+
126
+ def push(
127
+ self, chunk: str, *, emit: bool = True
128
+ ) -> tuple[SegmentationSnapshot, ...]:
129
+ if not isinstance(chunk, str):
130
+ raise TypeError("Parser chunks must be strings.")
131
+ self._raw_output_parts.append(chunk)
132
+
133
+ changed = False
134
+ segments = chunk.split("\n")
135
+ for segment in segments[:-1]:
136
+ self._append_buffer(segment)
137
+ raw = "".join(self._buffer_parts)
138
+ if raw.endswith("\r"):
139
+ raw = raw[:-1]
140
+ self._clear_buffer()
141
+ prior = (len(self._records), len(self._diagnostics))
142
+ self._accept_line(raw)
143
+ changed = changed or prior != (
144
+ len(self._records),
145
+ len(self._diagnostics),
146
+ )
147
+ self._append_buffer(segments[-1])
148
+
149
+ if changed:
150
+ self._revision += 1
151
+ return (self._snapshot(),) if changed and emit is not False else ()
152
+
153
+ def finish(
154
+ self, outcome: ResponseStreamOutcome
155
+ ) -> ParserFinish[SegmentationSnapshot, SegmentationResult]:
156
+ events: tuple[SegmentationSnapshot, ...] = ()
157
+ if self._buffer_parts:
158
+ raw = "".join(self._buffer_parts)
159
+ if raw.endswith("\r"):
160
+ raw = raw[:-1]
161
+ self._clear_buffer()
162
+ prior = (len(self._records), len(self._diagnostics))
163
+ self._accept_line(raw)
164
+ if prior != (len(self._records), len(self._diagnostics)):
165
+ self._revision += 1
166
+ events = (self._snapshot(),)
167
+ return ParserFinish(events=events, result=self._result(outcome))
168
+
169
+ def _append_buffer(self, value: str) -> None:
170
+ if value:
171
+ self._buffer_parts.append(value)
172
+
173
+ def _clear_buffer(self) -> None:
174
+ self._buffer_parts.clear()
175
+
176
+ def _accept_line(self, raw: str) -> bool:
177
+ self._line += 1
178
+ line = _trim(raw)
179
+ if not line:
180
+ return False
181
+
182
+ if line.startswith("<"):
183
+ accepted = self._accept_compact_api_line(line, raw)
184
+ if accepted is not None:
185
+ return accepted
186
+
187
+ match = _RECORD_PATTERN.fullmatch(line)
188
+ if match is None:
189
+ if line.startswith("[frame=") or line.startswith("object="):
190
+ self._diagnose("malformed_record", "Malformed structured record.", raw)
191
+ else:
192
+ self._add_record(
193
+ SegmentationTextRecord(order=len(self._records), text=raw)
194
+ )
195
+ return True
196
+
197
+ frame_index = None if match.group(1) is None else _integer_text(match.group(1))
198
+ if match.group(1) is not None and (frame_index is None or frame_index < 0):
199
+ self._diagnose(
200
+ "invalid_frame",
201
+ "Frame references must be non-negative safe integers.",
202
+ raw,
203
+ )
204
+ return True
205
+ if self._media == "image" and frame_index is not None:
206
+ self._diagnose(
207
+ "unexpected_frame",
208
+ "Image segmentation records cannot include a frame reference.",
209
+ raw,
210
+ )
211
+ return True
212
+
213
+ object_id = match.group(2)
214
+ kind = match.group(3)
215
+ value = match.group(4)
216
+ frame = None if frame_index is None else FrameReference(frame_index)
217
+ if kind == "point":
218
+ self._accept_point(object_id, frame, value, raw)
219
+ elif kind == "box":
220
+ self._accept_box(object_id, frame, value, raw)
221
+ else:
222
+ self._accept_mask(object_id, frame, value, raw)
223
+ return True
224
+
225
+ def _accept_compact_api_line(self, line: str, raw: str) -> bool | None:
226
+ header = _COMPACT_API_HEADER_PATTERN.fullmatch(line)
227
+ if header is None:
228
+ return None
229
+ frame_index = _integer_text(header.group(1))
230
+ if frame_index is None or frame_index < 0:
231
+ self._diagnose(
232
+ "invalid_frame",
233
+ "Frame references must be non-negative safe integers.",
234
+ raw,
235
+ )
236
+ return True
237
+ if self._media == "image" and frame_index != 0:
238
+ self._diagnose(
239
+ "unexpected_frame",
240
+ "Image segmentation records require frame zero.",
241
+ raw,
242
+ )
243
+ return True
244
+
245
+ frame = FrameReference(frame_index) if self._media == "video" else None
246
+ remainder = header.group(2)
247
+ if not remainder:
248
+ self._diagnose(
249
+ "malformed_record",
250
+ "Malformed compact SAM API object record.",
251
+ raw,
252
+ )
253
+ return True
254
+ accepted = False
255
+ while remainder:
256
+ match = _COMPACT_API_OBJECT_PATTERN.match(remainder)
257
+ if match is None:
258
+ self._diagnose(
259
+ "malformed_record",
260
+ "Malformed compact SAM API object record.",
261
+ raw,
262
+ )
263
+ return True
264
+ object_id = match.group(1)
265
+ coordinate_values = tuple(
266
+ _integer_text(value) for value in match.group(2, 3, 4, 5, 6, 7)
267
+ )
268
+ if any(value is None for value in coordinate_values):
269
+ self._diagnose(
270
+ "invalid_box", "Canonical SAM box coordinates are invalid.", raw
271
+ )
272
+ return True
273
+ (
274
+ left,
275
+ top,
276
+ inclusive_right,
277
+ inclusive_bottom,
278
+ source_width,
279
+ source_height,
280
+ ) = cast(tuple[int, int, int, int, int, int], coordinate_values)
281
+ if (
282
+ source_width <= 0
283
+ or source_height <= 0
284
+ or left < 0
285
+ or top < 0
286
+ or inclusive_right < left
287
+ or inclusive_bottom < top
288
+ or inclusive_right >= source_width
289
+ or inclusive_bottom >= source_height
290
+ ):
291
+ self._diagnose(
292
+ "invalid_box", "Canonical SAM box coordinates are invalid.", raw
293
+ )
294
+ return True
295
+ bounds = SegmentationMaskBounds(
296
+ left=left,
297
+ top=top,
298
+ right=inclusive_right + 1,
299
+ bottom=inclusive_bottom + 1,
300
+ )
301
+ self._add_record(
302
+ SegmentationBoxRecord(
303
+ order=len(self._records),
304
+ object_id=object_id,
305
+ frame=frame,
306
+ left=left,
307
+ top=top,
308
+ right=inclusive_right + 1,
309
+ bottom=inclusive_bottom + 1,
310
+ )
311
+ )
312
+ mask_height = match.group(8)
313
+ mask_width = match.group(9)
314
+ payload = match.group(10)
315
+ encoding = "lossless" if payload.startswith("~") else "one_bit"
316
+ self._accept_mask(
317
+ object_id,
318
+ frame,
319
+ f"{encoding};size={mask_width}x{mask_height};data={payload}",
320
+ raw,
321
+ bounds,
322
+ )
323
+ accepted = True
324
+ remainder = remainder[match.end() :]
325
+ return accepted
326
+
327
+ def _accept_point(
328
+ self,
329
+ object_id: str,
330
+ frame: FrameReference | None,
331
+ value: str,
332
+ raw: str,
333
+ ) -> None:
334
+ match = _POINT_PATTERN.fullmatch(_trim(value))
335
+ x = None if match is None else _javascript_number(match.group(1))
336
+ y = None if match is None else _javascript_number(match.group(2))
337
+ if x is None or y is None:
338
+ self._diagnose(
339
+ "malformed_point", "Point coordinates must be finite numbers.", raw
340
+ )
341
+ return
342
+ self._add_record(
343
+ SegmentationPointRecord(
344
+ order=len(self._records),
345
+ object_id=object_id,
346
+ frame=frame,
347
+ x=x,
348
+ y=y,
349
+ )
350
+ )
351
+
352
+ def _accept_box(
353
+ self,
354
+ object_id: str,
355
+ frame: FrameReference | None,
356
+ value: str,
357
+ raw: str,
358
+ ) -> None:
359
+ match = _BOX_PATTERN.fullmatch(_trim(value))
360
+ coordinates = (
361
+ None
362
+ if match is None
363
+ else tuple(_javascript_number(item) for item in match.group(1, 2, 3, 4))
364
+ )
365
+ if coordinates is None or any(item is None for item in coordinates):
366
+ self._diagnose(
367
+ "malformed_box", "Box coordinates must be finite numbers.", raw
368
+ )
369
+ return
370
+ left, top, right, bottom = cast(tuple[float, float, float, float], coordinates)
371
+ if right < left or bottom < top:
372
+ self._diagnose("invalid_box", "Box coordinates must be ordered.", raw)
373
+ return
374
+ self._add_record(
375
+ SegmentationBoxRecord(
376
+ order=len(self._records),
377
+ object_id=object_id,
378
+ frame=frame,
379
+ left=left,
380
+ top=top,
381
+ right=right,
382
+ bottom=bottom,
383
+ )
384
+ )
385
+
386
+ def _accept_mask(
387
+ self,
388
+ object_id: str,
389
+ frame: FrameReference | None,
390
+ value: str,
391
+ raw: str,
392
+ bounds: SegmentationMaskBounds | None = None,
393
+ ) -> None:
394
+ match = _MASK_PATTERN.fullmatch(_trim(value))
395
+ if match is None:
396
+ self._diagnose(
397
+ "malformed_mask",
398
+ "Mask records require encoding, size, and one complete payload.",
399
+ raw,
400
+ )
401
+ return
402
+ encoding = match.group(1)
403
+ width_number = _javascript_number(match.group(2))
404
+ height_number = _javascript_number(match.group(3))
405
+ width = int(width_number) if width_number is not None else 0
406
+ height = int(height_number) if height_number is not None else 0
407
+ payload = match.group(4)
408
+ if encoding not in ("one_bit", "lossless"):
409
+ self._diagnose(
410
+ "unsupported_mask_encoding",
411
+ f"Unsupported complete mask encoding: {encoding}.",
412
+ raw,
413
+ )
414
+ return
415
+ area = float(width) * float(height)
416
+ if (
417
+ width <= 0
418
+ or height <= 0
419
+ or not math.isfinite(area)
420
+ or not area.is_integer()
421
+ or abs(area) > _MAXIMUM_SAFE_INTEGER
422
+ ):
423
+ self._diagnose(
424
+ "invalid_mask_size", "Mask dimensions must be positive.", raw
425
+ )
426
+ return
427
+ mask = SegmentationMask(
428
+ encoding=encoding, payload=payload, width=width, height=height
429
+ )
430
+ try:
431
+ decode_mask_to_raster(mask)
432
+ except InvalidSegmentationMaskError as error:
433
+ self._diagnose("invalid_mask_payload", str(error), raw)
434
+ return
435
+
436
+ identity = SegmentationMaskIdentity(
437
+ media=self._media,
438
+ frame_index=None if frame is None else frame.frame_index,
439
+ object_id=object_id,
440
+ )
441
+ revision = self._revisions.get(identity, 0) + 1
442
+ self._revisions[identity] = revision
443
+ self._add_record(
444
+ SegmentationMaskRecord(
445
+ order=len(self._records),
446
+ object_id=object_id,
447
+ frame=frame,
448
+ identity=identity,
449
+ revision=revision,
450
+ mask=mask,
451
+ bounds=bounds,
452
+ )
453
+ )
454
+
455
+ def _add_record(self, record: SegmentationRecord) -> None:
456
+ self._records.append(record)
457
+
458
+ def _diagnose(self, code: str, message: str, raw: str) -> None:
459
+ self._diagnostics.append(
460
+ SegmentationDiagnostic(
461
+ severity="error",
462
+ code=code,
463
+ message=message,
464
+ line=self._line,
465
+ raw=raw,
466
+ )
467
+ )
468
+
469
+ def _snapshot(self) -> SegmentationSnapshot:
470
+ records = tuple(self._records)
471
+ diagnostics = tuple(self._diagnostics)
472
+ if self._media == "image":
473
+ return ImageSegmentationSnapshot(
474
+ revision=self._revision,
475
+ records=records,
476
+ diagnostics=diagnostics,
477
+ raw_output="".join(self._raw_output_parts),
478
+ )
479
+ return VideoSegmentationSnapshot(
480
+ revision=self._revision,
481
+ records=records,
482
+ diagnostics=diagnostics,
483
+ raw_output="".join(self._raw_output_parts),
484
+ )
485
+
486
+ def _result(self, outcome: ResponseStreamOutcome) -> SegmentationResult:
487
+ records = tuple(self._records)
488
+ diagnostics = tuple(self._diagnostics)
489
+ if self._media == "image":
490
+ return ImageSegmentationResult(
491
+ revision=self._revision,
492
+ records=records,
493
+ diagnostics=diagnostics,
494
+ raw_output="".join(self._raw_output_parts),
495
+ outcome=outcome,
496
+ )
497
+ return VideoSegmentationResult(
498
+ revision=self._revision,
499
+ records=records,
500
+ diagnostics=diagnostics,
501
+ raw_output="".join(self._raw_output_parts),
502
+ outcome=outcome,
503
+ )
504
+
505
+
506
+ @dataclass(frozen=True, slots=True)
507
+ class _SegmentationFormat:
508
+ _media: SegmentationMedia
509
+
510
+ def create_parser(self) -> _SegmentationParser:
511
+ return _SegmentationParser(self._media)
512
+
513
+
514
+ def image_segmentation_format() -> ResponseFormat[
515
+ ImageSegmentationSnapshot, ImageSegmentationResult
516
+ ]:
517
+ """Create an image format whose parsers have isolated incremental state."""
518
+
519
+ return cast(
520
+ ResponseFormat[ImageSegmentationSnapshot, ImageSegmentationResult],
521
+ _SegmentationFormat("image"),
522
+ )
523
+
524
+
525
+ def video_segmentation_format() -> ResponseFormat[
526
+ VideoSegmentationSnapshot, VideoSegmentationResult
527
+ ]:
528
+ """Create a video format whose parsers have isolated incremental state."""
529
+
530
+ return cast(
531
+ ResponseFormat[VideoSegmentationSnapshot, VideoSegmentationResult],
532
+ _SegmentationFormat("video"),
533
+ )