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,539 @@
1
+ # Copyright (c) Meta Platforms, Inc. and affiliates. All Rights Reserved.
2
+
3
+ from __future__ import annotations
4
+
5
+ import math
6
+ from array import array
7
+ from collections.abc import Sequence
8
+
9
+ from ._errors import InvalidSegmentationMaskError
10
+ from ._types import SegmentationMask
11
+
12
+ _EXCLUDED = frozenset({'"', "\\", ",", ";", "<", ">", "|"})
13
+ _ALPHABET = "".join(
14
+ character for code in range(0x21, 0x7F) if (character := chr(code)) not in _EXCLUDED
15
+ )
16
+ _RADIX = 85
17
+ _PREFIX_LENGTH = 5
18
+ _CHARACTER_VALUES = [-1] * 128
19
+ for _index, _character in enumerate(_ALPHABET):
20
+ _CHARACTER_VALUES[ord(_character)] = _index
21
+
22
+ _TOP = 1 << 24
23
+ _BOTTOM = 1 << 16
24
+ _MASK_8 = 0xFF
25
+ _MASK_16 = 0xFFFF
26
+ _MASK_32 = 0xFFFFFFFF
27
+ _INCREMENT = 14
28
+ _COUNT_LIMIT = 4096
29
+ _CONTEXT_COUNT = 1 << 12
30
+ _PADDING = 2
31
+ _SYMBOLS = 256
32
+ _TOTAL_INDEX = _SYMBOLS
33
+ _SPATIAL_MODES = 5
34
+ _ZERO_INCREMENT = 32
35
+ _ZERO_LIMIT = 16384
36
+ _ORDER_ONE_INCREMENT = 56
37
+ _ORDER_ONE_LIMIT = 8192
38
+ _ORDER_TWO_STEP = 14 * 44
39
+ _ORDER_TWO_LIMIT = 14 * 3584
40
+ _ORDER_ZERO_STEP = 2 * 16
41
+ _ORDER_ZERO_LIMIT = 2 * 2048
42
+ _ORDER_ZERO_INITIAL = 2
43
+ _ORDER_ONE_INITIAL = 1
44
+ _MAXIMUM_SAFE_INTEGER = (1 << 53) - 1
45
+
46
+
47
+ def _invalid(
48
+ message: str, cause: BaseException | None = None
49
+ ) -> InvalidSegmentationMaskError:
50
+ return InvalidSegmentationMaskError(message, cause=cause)
51
+
52
+
53
+ def _u8(value: int) -> int:
54
+ return value & _MASK_8
55
+
56
+
57
+ def _u16(value: int) -> int:
58
+ return value & _MASK_16
59
+
60
+
61
+ def _u32(value: int) -> int:
62
+ return value & _MASK_32
63
+
64
+
65
+ def _i32(value: int) -> int:
66
+ value &= _MASK_32
67
+ return value if value < 1 << 31 else value - (1 << 32)
68
+
69
+
70
+ def _as_safe_integer(value: object) -> int | None:
71
+ if isinstance(value, bool):
72
+ return None
73
+ if isinstance(value, int):
74
+ return value if abs(value) <= _MAXIMUM_SAFE_INTEGER else None
75
+ if isinstance(value, float) and math.isfinite(value) and value.is_integer():
76
+ integer = int(value)
77
+ return integer if abs(integer) <= _MAXIMUM_SAFE_INTEGER else None
78
+ return None
79
+
80
+
81
+ def _javascript_string_length(value: str) -> int:
82
+ return sum(2 if ord(character) > 0xFFFF else 1 for character in value)
83
+
84
+
85
+ def _digit(character: str) -> int:
86
+ code = ord(character)
87
+ value = _CHARACTER_VALUES[code] if code < len(_CHARACTER_VALUES) else -1
88
+ if value < 0 or value >= _RADIX:
89
+ raise _invalid("Mask payload contains an invalid character.")
90
+ return value
91
+
92
+
93
+ def _pack(input_: Sequence[int]) -> str:
94
+ length = len(input_)
95
+ remaining_length = length
96
+ prefix = [""] * _PREFIX_LENGTH
97
+ for index in range(_PREFIX_LENGTH - 1, -1, -1):
98
+ prefix[index] = _ALPHABET[remaining_length % _RADIX]
99
+ remaining_length //= _RADIX
100
+ if remaining_length != 0:
101
+ raise _invalid("Mask payload is too large.")
102
+
103
+ output = prefix
104
+ offset = 0
105
+ full = length - (length % 4)
106
+ while offset < full:
107
+ value = _u32(
108
+ (input_[offset] << 24)
109
+ | (input_[offset + 1] << 16)
110
+ | (input_[offset + 2] << 8)
111
+ | input_[offset + 3]
112
+ )
113
+ digits = [0] * 5
114
+ for index in range(4, -1, -1):
115
+ digits[index] = value % _RADIX
116
+ value = (value - digits[index]) // _RADIX
117
+ output.extend(_ALPHABET[encoded] for encoded in digits)
118
+ offset += 4
119
+
120
+ remainder = length - offset
121
+ if remainder > 0:
122
+ value = 0
123
+ for index in range(4):
124
+ value = _u32(
125
+ (value << 8) | (input_[offset + index] if index < remainder else 0)
126
+ )
127
+ digits = [0] * 5
128
+ for index in range(4, -1, -1):
129
+ digits[index] = value % _RADIX
130
+ value = (value - digits[index]) // _RADIX
131
+ output.extend(_ALPHABET[digits[index]] for index in range(remainder + 1))
132
+ return "".join(output)
133
+
134
+
135
+ def _unpack(payload: str) -> bytes:
136
+ payload_length = _javascript_string_length(payload)
137
+ if payload_length < _PREFIX_LENGTH:
138
+ raise _invalid("Mask payload is missing its length prefix.")
139
+ length = 0
140
+ for character in payload[:_PREFIX_LENGTH]:
141
+ length = length * _RADIX + _digit(character)
142
+ if length > _MAXIMUM_SAFE_INTEGER:
143
+ raise _invalid("Mask payload length is unsupported.")
144
+
145
+ remainder = length % 4
146
+ expected = (
147
+ _PREFIX_LENGTH + (length // 4) * 5 + (0 if remainder == 0 else remainder + 1)
148
+ )
149
+ if payload_length != expected:
150
+ raise _invalid("Mask payload length does not match its prefix.")
151
+
152
+ output = bytearray(length)
153
+ source = _PREFIX_LENGTH
154
+ destination = 0
155
+ full = length - remainder
156
+ while destination < full:
157
+ value = 0
158
+ for index in range(5):
159
+ value = value * _RADIX + _digit(payload[source + index])
160
+ if value > _MASK_32:
161
+ raise _invalid("Mask payload contains an out-of-range group.")
162
+ source += 5
163
+ output[destination] = (value >> 24) & _MASK_8
164
+ output[destination + 1] = (value >> 16) & _MASK_8
165
+ output[destination + 2] = (value >> 8) & _MASK_8
166
+ output[destination + 3] = value & _MASK_8
167
+ destination += 4
168
+
169
+ if remainder > 0:
170
+ value = 0
171
+ for index in range(5):
172
+ value = value * _RADIX + (
173
+ _digit(payload[source + index]) if index < remainder + 1 else _RADIX - 1
174
+ )
175
+ if value > _MASK_32:
176
+ raise _invalid("Mask payload contains an out-of-range tail.")
177
+ for index in range(remainder):
178
+ output[destination + index] = (value >> (24 - index * 8)) & _MASK_8
179
+
180
+ if _pack(output) != payload:
181
+ raise _invalid("Mask payload is not in canonical form.")
182
+ return bytes(output)
183
+
184
+
185
+ class _RangeEncoder:
186
+ __slots__ = ("_low", "_output", "_range")
187
+
188
+ def __init__(self) -> None:
189
+ self._low = 0
190
+ self._range = _MASK_32
191
+ self._output = bytearray()
192
+
193
+ def encode(self, cumulative: int, frequency: int, total: int) -> None:
194
+ scaled = self._range // total
195
+ self._low = _u32(self._low + scaled * cumulative)
196
+ self._range = scaled * frequency
197
+ while True:
198
+ if _u32(self._low ^ _u32(self._low + self._range)) < _TOP:
199
+ pass
200
+ elif self._range < _BOTTOM:
201
+ self._range = _u32(-self._low) & (_BOTTOM - 1)
202
+ else:
203
+ break
204
+ self._output.append((self._low >> 24) & _MASK_8)
205
+ self._low = _u32(self._low << 8)
206
+ self._range = _u32(self._range << 8)
207
+
208
+ def finish(self) -> bytes:
209
+ for _ in range(4):
210
+ self._output.append((self._low >> 24) & _MASK_8)
211
+ self._low = _u32(self._low << 8)
212
+ return bytes(self._output)
213
+
214
+
215
+ class _RangeDecoder:
216
+ __slots__ = ("_code", "_input", "_low", "_position", "_range", "_scaled")
217
+
218
+ def __init__(self, input_: bytes) -> None:
219
+ if len(input_) < 4:
220
+ raise _invalid("Mask payload is missing its finalization.")
221
+ self._input = input_
222
+ self._position = 0
223
+ self._low = 0
224
+ self._range = _MASK_32
225
+ self._code = 0
226
+ self._scaled = 0
227
+ for _ in range(4):
228
+ self._code = _u32((self._code << 8) | self._read())
229
+
230
+ def frequency(self, total: int) -> int:
231
+ self._scaled = self._range // total
232
+ if self._scaled == 0:
233
+ raise _invalid("Mask payload has invalid coding state.")
234
+ value = _u32(self._code - self._low) // self._scaled
235
+ return total - 1 if value >= total else value
236
+
237
+ def get_freq(self, total: int) -> int:
238
+ return self.frequency(total)
239
+
240
+ def decode(self, cumulative: int, frequency: int) -> None:
241
+ self._low = _u32(self._low + self._scaled * cumulative)
242
+ self._range = self._scaled * frequency
243
+ while True:
244
+ if _u32(self._low ^ _u32(self._low + self._range)) < _TOP:
245
+ pass
246
+ elif self._range < _BOTTOM:
247
+ self._range = _u32(-self._low) & (_BOTTOM - 1)
248
+ else:
249
+ break
250
+ self._code = _u32((self._code << 8) | self._read())
251
+ self._low = _u32(self._low << 8)
252
+ self._range = _u32(self._range << 8)
253
+
254
+ def _read(self) -> int:
255
+ if self._position >= len(self._input):
256
+ raise _invalid("Mask payload ended before decoding completed.")
257
+ value = self._input[self._position]
258
+ self._position += 1
259
+ return value
260
+
261
+
262
+ def _get_order_two(counts_by_context: dict[int, array[int]], key: int) -> array[int]:
263
+ counts = counts_by_context.get(key)
264
+ if counts is None:
265
+ counts = array("H", [0]) * (_SYMBOLS + 1)
266
+ counts_by_context[key] = counts
267
+ return counts
268
+
269
+
270
+ def _paeth(left: int, up: int, upper_left: int) -> int:
271
+ estimate = left + up - upper_left
272
+ left_distance = abs(estimate - left)
273
+ up_distance = abs(estimate - up)
274
+ upper_left_distance = abs(estimate - upper_left)
275
+ if left_distance <= up_distance and left_distance <= upper_left_distance:
276
+ return left
277
+ return up if up_distance <= upper_left_distance else upper_left
278
+
279
+
280
+ def _predict_byte(
281
+ output: bytearray,
282
+ offset: int,
283
+ x: int,
284
+ y: int,
285
+ mode: int,
286
+ width: int,
287
+ ) -> int:
288
+ if mode < 0 or mode >= _SPATIAL_MODES:
289
+ raise _invalid("Lossless mask payload uses an invalid predictor.")
290
+ left = output[offset - 1] if x > 0 else 0
291
+ up = output[offset - width] if y > 0 else 0
292
+ if mode == 1:
293
+ return up
294
+ if mode == 2:
295
+ return left
296
+ if mode == 3:
297
+ return (left + up) >> 1
298
+ upper_left = output[offset - width - 1] if x > 0 and y > 0 else 0
299
+ if mode == 0:
300
+ return _paeth(left, up, upper_left)
301
+ if upper_left >= max(left, up):
302
+ return min(left, up)
303
+ if upper_left <= min(left, up):
304
+ return max(left, up)
305
+ return left + up - upper_left
306
+
307
+
308
+ def _unzig(value: int, prediction: int) -> int:
309
+ delta = -((value + 1) >> 1) if value & 1 else value >> 1
310
+ return _u8(prediction + delta)
311
+
312
+
313
+ def _zero_context(residuals: bytearray, offset: int, x: int, y: int, width: int) -> int:
314
+ left = int(residuals[offset - 1] == 0) if x > 0 else 1
315
+ up = int(residuals[offset - width] == 0) if y > 0 else 1
316
+ upper_left = int(residuals[offset - width - 1] == 0) if x > 0 and y > 0 else 1
317
+ upper_right = (
318
+ int(residuals[offset - width + 1] == 0) if y > 0 and x < width - 1 else 1
319
+ )
320
+ return left | (up << 1) | (upper_left << 2) | (upper_right << 3)
321
+
322
+
323
+ def _decode_lossless_raster(payload: str, width: int, height: int) -> bytes:
324
+ if not payload.startswith("~"):
325
+ raise _invalid("lossless mask payloads must start with ~.")
326
+ packed = _unpack(payload[1:])
327
+ if len(packed) < 5:
328
+ raise _invalid("Lossless mask payload is truncated.")
329
+ selector = packed[0]
330
+ decoder = _RangeDecoder(packed[1:])
331
+ length = width * height
332
+ output = bytearray(length)
333
+ residuals = bytearray(length)
334
+ zero_counts = array("I", [1]) * 32
335
+ order_one = array("H", [_ORDER_ONE_INITIAL]) * (_SYMBOLS * _SYMBOLS)
336
+ order_one_totals = array("i", [_SYMBOLS * _ORDER_ONE_INITIAL]) * _SYMBOLS
337
+ order_zero = array("H", [_ORDER_ZERO_INITIAL]) * _SYMBOLS
338
+ order_zero_total = _SYMBOLS * _ORDER_ZERO_INITIAL
339
+ order_two: dict[int, array[int]] = {}
340
+ previous = 0
341
+
342
+ for offset in range(length):
343
+ x = offset % width
344
+ y = offset // width
345
+ zero_offset = _zero_context(residuals, offset, x, y, width) << 1
346
+ zero = zero_counts[zero_offset]
347
+ nonzero = zero_counts[zero_offset + 1]
348
+ zero_frequency = decoder.get_freq(zero + nonzero)
349
+ bit = 0 if zero_frequency < zero else 1
350
+ decoder.decode(0 if bit == 0 else zero, zero if bit == 0 else nonzero)
351
+ zero_counts[zero_offset + bit] = _u32(
352
+ zero_counts[zero_offset + bit] + _ZERO_INCREMENT
353
+ )
354
+ if zero_counts[zero_offset] + zero_counts[zero_offset + 1] >= _ZERO_LIMIT:
355
+ zero_counts[zero_offset] = (zero_counts[zero_offset] >> 1) or 1
356
+ zero_counts[zero_offset + 1] = (zero_counts[zero_offset + 1] >> 1) or 1
357
+
358
+ symbol = 0
359
+ if bit == 1:
360
+ above = residuals[offset - width] if offset >= width else 0
361
+ second = _get_order_two(order_two, (previous << 8) | above)
362
+ first_offset = previous << 8
363
+ excluded_zero = second[0] + order_one[first_offset] + order_zero[0]
364
+ total = (
365
+ second[_TOTAL_INDEX]
366
+ + order_one_totals[previous]
367
+ + order_zero_total
368
+ - excluded_zero
369
+ )
370
+ target = decoder.get_freq(total)
371
+ cumulative = 0
372
+ symbol = 1
373
+ frequency = second[1] + order_one[first_offset + 1] + order_zero[1]
374
+ while cumulative + frequency <= target:
375
+ cumulative += frequency
376
+ symbol += 1
377
+ if symbol >= _SYMBOLS:
378
+ raise _invalid("Lossless mask payload is malformed.")
379
+ frequency = (
380
+ second[symbol]
381
+ + order_one[first_offset + symbol]
382
+ + order_zero[symbol]
383
+ )
384
+ decoder.decode(cumulative, frequency)
385
+ second[symbol] = _u16(second[symbol] + _ORDER_TWO_STEP)
386
+ second[_TOTAL_INDEX] = _u16(second[_TOTAL_INDEX] + _ORDER_TWO_STEP)
387
+ if second[_TOTAL_INDEX] >= _ORDER_TWO_LIMIT:
388
+ total_after_scaling = 0
389
+ for value in range(_SYMBOLS):
390
+ second[value] = second[value] >> 1
391
+ total_after_scaling += second[value]
392
+ second[_TOTAL_INDEX] = _u16(total_after_scaling)
393
+ order_one[first_offset + symbol] = _u16(
394
+ order_one[first_offset + symbol] + _ORDER_ONE_INCREMENT
395
+ )
396
+ order_one_totals[previous] = _i32(
397
+ order_one_totals[previous] + _ORDER_ONE_INCREMENT
398
+ )
399
+ if order_one_totals[previous] >= _ORDER_ONE_LIMIT:
400
+ total_after_scaling = 0
401
+ for value in range(_SYMBOLS):
402
+ scaled = (order_one[first_offset + value] >> 1) or 1
403
+ order_one[first_offset + value] = scaled
404
+ total_after_scaling += scaled
405
+ order_one_totals[previous] = _i32(total_after_scaling)
406
+ order_zero[symbol] = _u16(order_zero[symbol] + _ORDER_ZERO_STEP)
407
+ order_zero_total += _ORDER_ZERO_STEP
408
+ if order_zero_total >= _ORDER_ZERO_LIMIT:
409
+ total_after_scaling = 0
410
+ for value in range(_SYMBOLS):
411
+ order_zero[value] = order_zero[value] >> 1
412
+ total_after_scaling += order_zero[value]
413
+ order_zero_total = total_after_scaling
414
+ residuals[offset] = symbol
415
+ output[offset] = _unzig(
416
+ symbol, _predict_byte(output, offset, x, y, selector, width)
417
+ )
418
+ previous = symbol
419
+ return bytes(output)
420
+
421
+
422
+ def _context_at(raster: bytearray, position: int, prior: int, earlier: int) -> int:
423
+ return (
424
+ raster[position - 1]
425
+ | (raster[position - 2] << 1)
426
+ | (raster[prior - 2] << 2)
427
+ | (raster[prior - 1] << 3)
428
+ | (raster[prior] << 4)
429
+ | (raster[prior + 1] << 5)
430
+ | (raster[prior + 2] << 6)
431
+ | (raster[earlier - 2] << 7)
432
+ | (raster[earlier - 1] << 8)
433
+ | (raster[earlier] << 9)
434
+ | (raster[earlier + 1] << 10)
435
+ | (raster[earlier + 2] << 11)
436
+ )
437
+
438
+
439
+ def _update_counts(counts: array[int], offset: int, bit: int) -> None:
440
+ counts[offset + bit] = _u16(counts[offset + bit] + _INCREMENT)
441
+ if counts[offset] + counts[offset + 1] >= _COUNT_LIMIT:
442
+ counts[offset] = (counts[offset] >> 1) or 1
443
+ counts[offset + 1] = (counts[offset + 1] >> 1) or 1
444
+
445
+
446
+ def _encode_raster(input_: Sequence[int], width: int, height: int) -> str:
447
+ encoder = _RangeEncoder()
448
+ counts = array("H", [1]) * (_CONTEXT_COUNT * 2)
449
+ padded_width = width + 2 * _PADDING
450
+ padded = bytearray(padded_width * (height + _PADDING))
451
+ for y in range(height):
452
+ row = (y + _PADDING) * padded_width + _PADDING
453
+ for x in range(width):
454
+ position = row + x
455
+ context = _context_at(
456
+ padded,
457
+ position,
458
+ position - padded_width,
459
+ position - 2 * padded_width,
460
+ )
461
+ offset = context << 1
462
+ zero = counts[offset]
463
+ one = counts[offset + 1]
464
+ bit = input_[y * width + x]
465
+ if bit not in (0, 1):
466
+ raise _invalid("Decoded mask is not binary.")
467
+ if bit == 0:
468
+ encoder.encode(0, zero, zero + one)
469
+ else:
470
+ encoder.encode(zero, one, zero + one)
471
+ padded[position] = bit
472
+ _update_counts(counts, offset, bit)
473
+ return f"!{_pack(encoder.finish())}"
474
+
475
+
476
+ def _decode_raster(payload: str, width: int, height: int) -> bytes:
477
+ decoder = _RangeDecoder(_unpack(payload[1:]))
478
+ counts = array("H", [1]) * (_CONTEXT_COUNT * 2)
479
+ padded_width = width + 2 * _PADDING
480
+ padded = bytearray(padded_width * (height + _PADDING))
481
+ output = bytearray(width * height)
482
+ for y in range(height):
483
+ row = (y + _PADDING) * padded_width + _PADDING
484
+ for x in range(width):
485
+ position = row + x
486
+ context = _context_at(
487
+ padded,
488
+ position,
489
+ position - padded_width,
490
+ position - 2 * padded_width,
491
+ )
492
+ offset = context << 1
493
+ zero = counts[offset]
494
+ one = counts[offset + 1]
495
+ value = decoder.frequency(zero + one)
496
+ bit = 0 if value < zero else 1
497
+ if bit == 0:
498
+ decoder.decode(0, zero)
499
+ else:
500
+ decoder.decode(zero, one)
501
+ padded[position] = bit
502
+ output[y * width + x] = bit
503
+ _update_counts(counts, offset, bit)
504
+ return bytes(output)
505
+
506
+
507
+ def decode_mask_to_raster(mask: SegmentationMask) -> bytes:
508
+ """Strictly decode a complete mask with structural checks and no quota ceiling."""
509
+
510
+ if mask.encoding not in ("one_bit", "lossless"):
511
+ raise _invalid(f"Unsupported complete mask encoding: {mask.encoding}.")
512
+ width = _as_safe_integer(mask.width)
513
+ height = _as_safe_integer(mask.height)
514
+ if width is None or height is None or width <= 0 or height <= 0:
515
+ raise _invalid("Mask dimensions must be positive safe integers.")
516
+ area = width * height
517
+ if area > _MAXIMUM_SAFE_INTEGER:
518
+ raise _invalid("Mask dimensions produce an unsafe decoded area.")
519
+ if not isinstance(mask.payload, str):
520
+ raise _invalid("Mask payload must be a string.")
521
+ if mask.encoding == "one_bit" and not mask.payload.startswith("!"):
522
+ raise _invalid("one_bit mask payloads must start with !.")
523
+ if mask.encoding == "lossless" and not mask.payload.startswith("~"):
524
+ raise _invalid("lossless mask payloads must start with ~.")
525
+
526
+ try:
527
+ if mask.encoding == "lossless":
528
+ coverage = _decode_lossless_raster(mask.payload, width, height)
529
+ return bytes(value >= 129 for value in coverage)
530
+ raster = _decode_raster(mask.payload, width, height)
531
+ if _encode_raster(raster, width, height) != mask.payload:
532
+ raise _invalid("Mask payload is not canonical or has invalid finalization.")
533
+ return raster
534
+ except InvalidSegmentationMaskError:
535
+ raise
536
+ except MemoryError:
537
+ raise
538
+ except Exception as error:
539
+ raise _invalid("Mask payload could not be decoded.", error) from error