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,367 @@
|
|
|
1
|
+
# Copyright (c) Meta Platforms, Inc. and affiliates. All Rights Reserved.
|
|
2
|
+
#
|
|
3
|
+
# Portions of this file are a dependency-free port of the COCO API mask module
|
|
4
|
+
# (`common/maskApi.c` and `PythonAPI/pycocotools/_mask.pyx`).
|
|
5
|
+
# Copyright (c) 2014, Piotr Dollar and Tsung-Yi Lin. All rights reserved.
|
|
6
|
+
# Licensed under the Simplified BSD License; see THIRD_PARTY_NOTICES.md in the
|
|
7
|
+
# repository root for the full license text.
|
|
8
|
+
|
|
9
|
+
from __future__ import annotations
|
|
10
|
+
|
|
11
|
+
from dataclasses import dataclass
|
|
12
|
+
|
|
13
|
+
from ._mask_codec import decode_mask_to_raster
|
|
14
|
+
from ._types import SegmentationMask
|
|
15
|
+
|
|
16
|
+
_MASK_32 = 0xFFFFFFFF
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
@dataclass(frozen=True, slots=True)
|
|
20
|
+
class RLEObject:
|
|
21
|
+
size: tuple[int, int]
|
|
22
|
+
counts: str
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
@dataclass(slots=True)
|
|
26
|
+
class _DataArray:
|
|
27
|
+
data: bytes
|
|
28
|
+
shape: tuple[int, ...]
|
|
29
|
+
|
|
30
|
+
def reshape(self, shape: tuple[int, ...]) -> _DataArray:
|
|
31
|
+
return _DataArray(self.data, shape)
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
@dataclass(slots=True)
|
|
35
|
+
class _RLE:
|
|
36
|
+
h: int
|
|
37
|
+
w: int
|
|
38
|
+
m: int
|
|
39
|
+
cnts: list[int]
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
class _RLEs:
|
|
43
|
+
__slots__ = ("R", "n")
|
|
44
|
+
|
|
45
|
+
def __init__(self, n: int) -> None:
|
|
46
|
+
self.R = [_RLE(0, 0, 0, [0]) for _ in range(n)]
|
|
47
|
+
self.n = n
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class _Masks:
|
|
51
|
+
__slots__ = ("h", "mask", "n", "w")
|
|
52
|
+
|
|
53
|
+
def __init__(self, h: int, w: int, n: int) -> None:
|
|
54
|
+
self.mask = bytearray(h * w * n)
|
|
55
|
+
self.h = h
|
|
56
|
+
self.w = w
|
|
57
|
+
self.n = n
|
|
58
|
+
|
|
59
|
+
def to_data_array(self) -> _DataArray:
|
|
60
|
+
return _DataArray(bytes(self.mask), (self.h, self.w, self.n))
|
|
61
|
+
|
|
62
|
+
|
|
63
|
+
def _rle_init(R: _RLE, h: int, w: int, m: int, cnts: list[int]) -> None:
|
|
64
|
+
R.h = h
|
|
65
|
+
R.w = w
|
|
66
|
+
R.m = m
|
|
67
|
+
R.cnts = [0] if m == 0 else cnts
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
def _encode(bitmask: _DataArray) -> RLEObject | list[RLEObject]:
|
|
71
|
+
if len(bitmask.shape) == 3:
|
|
72
|
+
return _encode_many(bitmask)
|
|
73
|
+
if len(bitmask.shape) == 2:
|
|
74
|
+
h = bitmask.shape[0]
|
|
75
|
+
w = bitmask.shape[1]
|
|
76
|
+
result = _encode_many(bitmask.reshape((h, w, 1)))
|
|
77
|
+
return result[0]
|
|
78
|
+
raise ValueError("wrong shape of bitmask")
|
|
79
|
+
|
|
80
|
+
|
|
81
|
+
def _encode_many(bitmask: _DataArray) -> list[RLEObject]:
|
|
82
|
+
h = bitmask.shape[0]
|
|
83
|
+
w = bitmask.shape[1]
|
|
84
|
+
n = bitmask.shape[2]
|
|
85
|
+
rles = _RLEs(n)
|
|
86
|
+
_rle_encode(rles.R, bitmask.data, h, w, n)
|
|
87
|
+
return _to_string(rles)
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
def _decode(rle_objects: RLEObject | list[RLEObject]) -> _DataArray:
|
|
91
|
+
rles = _from_string(rle_objects)
|
|
92
|
+
h = rles.R[0].h
|
|
93
|
+
w = rles.R[0].w
|
|
94
|
+
n = rles.n
|
|
95
|
+
masks = _Masks(h, w, n)
|
|
96
|
+
_rle_decode(rles.R, masks.mask, n)
|
|
97
|
+
data_array = masks.to_data_array()
|
|
98
|
+
return data_array if isinstance(rle_objects, list) else data_array.reshape((h, w))
|
|
99
|
+
|
|
100
|
+
|
|
101
|
+
def _rle_encode(R: list[_RLE], M: bytes, h: int, w: int, n: int) -> None:
|
|
102
|
+
area = w * h
|
|
103
|
+
cnts: list[int] = []
|
|
104
|
+
for i in range(n):
|
|
105
|
+
from_ = area * i
|
|
106
|
+
to = area * (i + 1)
|
|
107
|
+
values = M[from_:to]
|
|
108
|
+
k = 0
|
|
109
|
+
previous = 0
|
|
110
|
+
count = 0
|
|
111
|
+
for value in values:
|
|
112
|
+
if value != previous:
|
|
113
|
+
if k == len(cnts):
|
|
114
|
+
cnts.append(count)
|
|
115
|
+
else:
|
|
116
|
+
cnts[k] = count
|
|
117
|
+
k += 1
|
|
118
|
+
count = 0
|
|
119
|
+
previous = value
|
|
120
|
+
count += 1
|
|
121
|
+
if k == len(cnts):
|
|
122
|
+
cnts.append(count)
|
|
123
|
+
else:
|
|
124
|
+
cnts[k] = count
|
|
125
|
+
k += 1
|
|
126
|
+
_rle_init(R[i], h, w, k, list(cnts))
|
|
127
|
+
|
|
128
|
+
|
|
129
|
+
def _rle_decode(R: list[_RLE], M: bytearray, n: int) -> None:
|
|
130
|
+
position = 0
|
|
131
|
+
for i in range(n):
|
|
132
|
+
value = False
|
|
133
|
+
for j in range(R[i].m):
|
|
134
|
+
for _ in range(R[i].cnts[j]):
|
|
135
|
+
M[position] = 0 if value is False else 1
|
|
136
|
+
position += 1
|
|
137
|
+
value = not value
|
|
138
|
+
|
|
139
|
+
|
|
140
|
+
def _i32(value: int) -> int:
|
|
141
|
+
value &= _MASK_32
|
|
142
|
+
return value if value < 1 << 31 else value - (1 << 32)
|
|
143
|
+
|
|
144
|
+
|
|
145
|
+
def _rle_to_string(R: _RLE) -> str:
|
|
146
|
+
output: list[str] = []
|
|
147
|
+
for i in range(R.m):
|
|
148
|
+
value = R.cnts[i]
|
|
149
|
+
if i > 2:
|
|
150
|
+
value -= R.cnts[i - 2]
|
|
151
|
+
value = _i32(value)
|
|
152
|
+
more = True
|
|
153
|
+
while more:
|
|
154
|
+
character = value & 0x1F
|
|
155
|
+
value >>= 5
|
|
156
|
+
more = value != -1 if character & 0x10 else value != 0
|
|
157
|
+
if more:
|
|
158
|
+
character |= 0x20
|
|
159
|
+
character += 48
|
|
160
|
+
output.append(chr(character))
|
|
161
|
+
return "".join(output)
|
|
162
|
+
|
|
163
|
+
|
|
164
|
+
def _to_string(rles: _RLEs) -> list[RLEObject]:
|
|
165
|
+
return [
|
|
166
|
+
RLEObject(size=(rle.h, rle.w), counts=_rle_to_string(rle)) for rle in rles.R
|
|
167
|
+
]
|
|
168
|
+
|
|
169
|
+
|
|
170
|
+
def _from_string(
|
|
171
|
+
input_rle_objects: RLEObject | list[RLEObject],
|
|
172
|
+
) -> _RLEs:
|
|
173
|
+
rle_objects = (
|
|
174
|
+
input_rle_objects
|
|
175
|
+
if isinstance(input_rle_objects, list)
|
|
176
|
+
else [input_rle_objects]
|
|
177
|
+
)
|
|
178
|
+
rles = _RLEs(len(rle_objects))
|
|
179
|
+
for index, rle_object in enumerate(rle_objects):
|
|
180
|
+
_rle_from_string(
|
|
181
|
+
rles.R[index], rle_object.counts, rle_object.size[0], rle_object.size[1]
|
|
182
|
+
)
|
|
183
|
+
return rles
|
|
184
|
+
|
|
185
|
+
|
|
186
|
+
def _rle_from_string(R: _RLE, value: str, h: int, w: int) -> None:
|
|
187
|
+
counts: list[int] = []
|
|
188
|
+
position = 0
|
|
189
|
+
while position < len(value):
|
|
190
|
+
count = 0
|
|
191
|
+
shift = 0
|
|
192
|
+
more = 1
|
|
193
|
+
while more:
|
|
194
|
+
character = ord(value[position]) - 48
|
|
195
|
+
count = _i32(count | _i32((character & 0x1F) << ((5 * shift) & 0x1F)))
|
|
196
|
+
more = character & 0x20
|
|
197
|
+
position += 1
|
|
198
|
+
shift += 1
|
|
199
|
+
if not more and character & 0x10:
|
|
200
|
+
count = _i32(count | _i32(-1 << ((5 * shift) & 0x1F)))
|
|
201
|
+
if len(counts) > 2:
|
|
202
|
+
count += counts[-2]
|
|
203
|
+
counts.append(count)
|
|
204
|
+
_rle_init(R, h, w, len(counts), counts)
|
|
205
|
+
|
|
206
|
+
|
|
207
|
+
_Point = tuple[float, float]
|
|
208
|
+
_Segment = tuple[float, float, float, float]
|
|
209
|
+
|
|
210
|
+
|
|
211
|
+
def _decode_to_svg_path(rle_objects: RLEObject | list[RLEObject]) -> list[str]:
|
|
212
|
+
rle_array = rle_objects if isinstance(rle_objects, list) else [rle_objects]
|
|
213
|
+
paths: list[str] = []
|
|
214
|
+
for rle_object in rle_array:
|
|
215
|
+
mask_data = _decode(rle_object)
|
|
216
|
+
height = rle_object.size[0]
|
|
217
|
+
width = rle_object.size[1]
|
|
218
|
+
paths.append(_trace_contours(mask_data.data, height, width))
|
|
219
|
+
return paths
|
|
220
|
+
|
|
221
|
+
|
|
222
|
+
def _trace_contours(mask: bytes, height: int, width: int) -> str:
|
|
223
|
+
def get_pixel(row: int, col: int) -> int:
|
|
224
|
+
if row < 0 or row >= height or col < 0 or col >= width:
|
|
225
|
+
return 0
|
|
226
|
+
return mask[col * height + row]
|
|
227
|
+
|
|
228
|
+
segments: list[_Segment] = []
|
|
229
|
+
for row in range(height + 1):
|
|
230
|
+
for col in range(width + 1):
|
|
231
|
+
tl = get_pixel(row - 1, col - 1)
|
|
232
|
+
tr = get_pixel(row - 1, col)
|
|
233
|
+
bl = get_pixel(row, col - 1)
|
|
234
|
+
br = get_pixel(row, col)
|
|
235
|
+
case_index = (tl << 3) | (tr << 2) | (br << 1) | bl
|
|
236
|
+
top = (col - 0.5, row - 1.0)
|
|
237
|
+
bottom = (col - 0.5, float(row))
|
|
238
|
+
left = (col - 1.0, row - 0.5)
|
|
239
|
+
right = (float(col), row - 0.5)
|
|
240
|
+
|
|
241
|
+
if case_index in (0, 15):
|
|
242
|
+
continue
|
|
243
|
+
if case_index in (1, 14):
|
|
244
|
+
segments.append((*left, *bottom))
|
|
245
|
+
elif case_index in (2, 13):
|
|
246
|
+
segments.append((*bottom, *right))
|
|
247
|
+
elif case_index in (3, 12):
|
|
248
|
+
segments.append((*left, *right))
|
|
249
|
+
elif case_index in (4, 11):
|
|
250
|
+
segments.append((*top, *right))
|
|
251
|
+
elif case_index == 5:
|
|
252
|
+
segments.append((*left, *top))
|
|
253
|
+
segments.append((*bottom, *right))
|
|
254
|
+
elif case_index in (6, 9):
|
|
255
|
+
segments.append((*top, *bottom))
|
|
256
|
+
elif case_index in (7, 8):
|
|
257
|
+
segments.append((*left, *top))
|
|
258
|
+
elif case_index == 10:
|
|
259
|
+
segments.append((*top, *right))
|
|
260
|
+
segments.append((*left, *bottom))
|
|
261
|
+
|
|
262
|
+
if not segments:
|
|
263
|
+
return ""
|
|
264
|
+
|
|
265
|
+
paths = _connect_segments(segments)
|
|
266
|
+
svg_paths: list[str] = []
|
|
267
|
+
for path in paths:
|
|
268
|
+
if len(path) < 2:
|
|
269
|
+
continue
|
|
270
|
+
commands = [f"M{_number(path[0][0])} {_number(path[0][1])}"]
|
|
271
|
+
commands.extend(
|
|
272
|
+
f"L{_number(point[0])} {_number(point[1])}" for point in path[1:]
|
|
273
|
+
)
|
|
274
|
+
commands.append("Z")
|
|
275
|
+
svg_paths.append("".join(commands))
|
|
276
|
+
return "".join(svg_paths)
|
|
277
|
+
|
|
278
|
+
|
|
279
|
+
def _number(value: float) -> str:
|
|
280
|
+
return str(int(value)) if value.is_integer() else str(value)
|
|
281
|
+
|
|
282
|
+
|
|
283
|
+
def _connect_segments(segments: list[_Segment]) -> list[list[_Point]]:
|
|
284
|
+
adjacency: dict[_Point, list[_Point]] = {}
|
|
285
|
+
insertion_order: list[_Point] = []
|
|
286
|
+
|
|
287
|
+
def add(point: _Point, neighbor: _Point) -> None:
|
|
288
|
+
if point not in adjacency:
|
|
289
|
+
adjacency[point] = []
|
|
290
|
+
insertion_order.append(point)
|
|
291
|
+
adjacency[point].append(neighbor)
|
|
292
|
+
|
|
293
|
+
for x1, y1, x2, y2 in segments:
|
|
294
|
+
add((x1, y1), (x2, y2))
|
|
295
|
+
add((x2, y2), (x1, y1))
|
|
296
|
+
|
|
297
|
+
paths: list[list[_Point]] = []
|
|
298
|
+
visited: set[_Point] = set()
|
|
299
|
+
for start in insertion_order:
|
|
300
|
+
neighbors = adjacency.get(start)
|
|
301
|
+
if start in visited or not neighbors:
|
|
302
|
+
continue
|
|
303
|
+
|
|
304
|
+
path: list[_Point] = []
|
|
305
|
+
current = start
|
|
306
|
+
previous: _Point | None = None
|
|
307
|
+
while True:
|
|
308
|
+
path.append(current)
|
|
309
|
+
current_neighbors = adjacency.get(current)
|
|
310
|
+
if not current_neighbors:
|
|
311
|
+
break
|
|
312
|
+
|
|
313
|
+
next_ = next(
|
|
314
|
+
(neighbor for neighbor in current_neighbors if neighbor != previous),
|
|
315
|
+
None,
|
|
316
|
+
)
|
|
317
|
+
if next_ is None:
|
|
318
|
+
break
|
|
319
|
+
|
|
320
|
+
filtered = [neighbor for neighbor in current_neighbors if neighbor != next_]
|
|
321
|
+
if filtered:
|
|
322
|
+
adjacency[current] = filtered
|
|
323
|
+
else:
|
|
324
|
+
del adjacency[current]
|
|
325
|
+
|
|
326
|
+
next_neighbors = adjacency.get(next_)
|
|
327
|
+
if next_neighbors is not None:
|
|
328
|
+
next_filtered = [
|
|
329
|
+
neighbor for neighbor in next_neighbors if neighbor != current
|
|
330
|
+
]
|
|
331
|
+
if next_filtered:
|
|
332
|
+
adjacency[next_] = next_filtered
|
|
333
|
+
else:
|
|
334
|
+
del adjacency[next_]
|
|
335
|
+
|
|
336
|
+
if next_ == start:
|
|
337
|
+
break
|
|
338
|
+
|
|
339
|
+
previous = current
|
|
340
|
+
current = next_
|
|
341
|
+
visited.add(previous)
|
|
342
|
+
|
|
343
|
+
if len(path) >= 3:
|
|
344
|
+
paths.append(path)
|
|
345
|
+
|
|
346
|
+
return paths
|
|
347
|
+
|
|
348
|
+
|
|
349
|
+
def decode_mask_to_rle(mask: SegmentationMask) -> RLEObject:
|
|
350
|
+
raster = decode_mask_to_raster(mask)
|
|
351
|
+
height = int(mask.height)
|
|
352
|
+
width = int(mask.width)
|
|
353
|
+
coco = bytearray(len(raster))
|
|
354
|
+
for row in range(height):
|
|
355
|
+
for col in range(width):
|
|
356
|
+
coco[col * height + row] = raster[row * width + col]
|
|
357
|
+
encoded = _encode(_DataArray(bytes(coco), (height, width)))
|
|
358
|
+
if isinstance(encoded, list):
|
|
359
|
+
raise RuntimeError("COCO encoding returned multiple masks for one raster.")
|
|
360
|
+
return encoded
|
|
361
|
+
|
|
362
|
+
|
|
363
|
+
def decode_mask_to_svg_path(mask: SegmentationMask) -> str:
|
|
364
|
+
paths = _decode_to_svg_path(decode_mask_to_rle(mask))
|
|
365
|
+
if len(paths) != 1:
|
|
366
|
+
raise RuntimeError("SVG conversion returned multiple paths for one mask.")
|
|
367
|
+
return paths[0]
|