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.
- meta_sam_parser/__init__.py +99 -0
- meta_sam_parser/_errors.py +116 -0
- meta_sam_parser/_mask_codec.py +539 -0
- meta_sam_parser/_mask_conversion.py +367 -0
- meta_sam_parser/_segmentation.py +533 -0
- meta_sam_parser/_stream.py +618 -0
- meta_sam_parser/_types.py +224 -0
- meta_sam_parser/py.typed +0 -0
- meta_sam_parser-0.0.2.dist-info/METADATA +394 -0
- meta_sam_parser-0.0.2.dist-info/RECORD +12 -0
- meta_sam_parser-0.0.2.dist-info/WHEEL +4 -0
- meta_sam_parser-0.0.2.dist-info/licenses/LICENSE +61 -0
|
@@ -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
|
+
)
|