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,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]