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