jxl 1.0.0

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.
Files changed (94) hide show
  1. checksums.yaml +7 -0
  2. data/.rubocop.yml +93 -0
  3. data/LICENSE.txt +21 -0
  4. data/README.md +90 -0
  5. data/Rakefile +31 -0
  6. data/docs/adr/0001-buffer-representation.md +12 -0
  7. data/docs/bench_history.md +11 -0
  8. data/docs/spec_map.md +16 -0
  9. data/exe/cjxl +21 -0
  10. data/exe/djxl +66 -0
  11. data/exe/jxlinfo +31 -0
  12. data/lib/jxl/basic_info.rb +16 -0
  13. data/lib/jxl/bit/field.rb +126 -0
  14. data/lib/jxl/bit/reader.rb +78 -0
  15. data/lib/jxl/bit/writer.rb +46 -0
  16. data/lib/jxl/codestream/reader.rb +68 -0
  17. data/lib/jxl/color/gamut.rb +80 -0
  18. data/lib/jxl/color/icc_profile.rb +243 -0
  19. data/lib/jxl/color/opsin.rb +28 -0
  20. data/lib/jxl/color/transfer.rb +108 -0
  21. data/lib/jxl/color/ycbcr.rb +63 -0
  22. data/lib/jxl/container/box.rb +7 -0
  23. data/lib/jxl/container/parser.rb +134 -0
  24. data/lib/jxl/container/signature.rb +19 -0
  25. data/lib/jxl/dct.rb +81 -0
  26. data/lib/jxl/decoder.rb +85 -0
  27. data/lib/jxl/encoder.rb +180 -0
  28. data/lib/jxl/entropy/ans_distribution.rb +212 -0
  29. data/lib/jxl/entropy/decoder.rb +226 -0
  30. data/lib/jxl/entropy/encoder.rb +69 -0
  31. data/lib/jxl/entropy/hybrid_uint.rb +55 -0
  32. data/lib/jxl/entropy/icc_stream.rb +293 -0
  33. data/lib/jxl/entropy/permutation.rb +64 -0
  34. data/lib/jxl/entropy/prefix_code.rb +177 -0
  35. data/lib/jxl/errors.rb +12 -0
  36. data/lib/jxl/features/noise.rb +140 -0
  37. data/lib/jxl/features/patches.rb +147 -0
  38. data/lib/jxl/features/splines.rb +206 -0
  39. data/lib/jxl/features/spot_colour.rb +29 -0
  40. data/lib/jxl/features/upsampling.rb +99 -0
  41. data/lib/jxl/features/upsampling_weights.bin +0 -0
  42. data/lib/jxl/filter/epf.rb +119 -0
  43. data/lib/jxl/filter/gaborish.rb +36 -0
  44. data/lib/jxl/frame/toc.rb +31 -0
  45. data/lib/jxl/headers/bit_depth.rb +32 -0
  46. data/lib/jxl/headers/colour_encoding.rb +81 -0
  47. data/lib/jxl/headers/custom_transform.rb +46 -0
  48. data/lib/jxl/headers/frame_header.rb +251 -0
  49. data/lib/jxl/headers/image_metadata.rb +128 -0
  50. data/lib/jxl/headers/size_header.rb +29 -0
  51. data/lib/jxl/image.rb +63 -0
  52. data/lib/jxl/io/npy.rb +82 -0
  53. data/lib/jxl/io/pam.rb +14 -0
  54. data/lib/jxl/io/pfm.rb +21 -0
  55. data/lib/jxl/io/pgx.rb +22 -0
  56. data/lib/jxl/io/png.rb +132 -0
  57. data/lib/jxl/io/ppm.rb +63 -0
  58. data/lib/jxl/modular/decoder.rb +371 -0
  59. data/lib/jxl/modular/group_header.rb +82 -0
  60. data/lib/jxl/modular/ma_tree.rb +95 -0
  61. data/lib/jxl/modular/predictor.rb +107 -0
  62. data/lib/jxl/modular/stream.rb +49 -0
  63. data/lib/jxl/modular/transform.rb +235 -0
  64. data/lib/jxl/modular/weighted.rb +85 -0
  65. data/lib/jxl/num.rb +37 -0
  66. data/lib/jxl/plane.rb +35 -0
  67. data/lib/jxl/render/blender.rb +61 -0
  68. data/lib/jxl/render/orientation.rb +42 -0
  69. data/lib/jxl/trace.rb +36 -0
  70. data/lib/jxl/vardct/ac_strategy.rb +65 -0
  71. data/lib/jxl/vardct/afv_basis.bin +0 -0
  72. data/lib/jxl/vardct/block_context_map.rb +63 -0
  73. data/lib/jxl/vardct/chroma_from_luma.rb +27 -0
  74. data/lib/jxl/vardct/coeff_order.rb +35 -0
  75. data/lib/jxl/vardct/decoder.rb +136 -0
  76. data/lib/jxl/vardct/default_dequant.bin.z +0 -0
  77. data/lib/jxl/vardct/dequant.rb +284 -0
  78. data/lib/jxl/vardct/hf_global.rb +30 -0
  79. data/lib/jxl/vardct/lf_global.rb +43 -0
  80. data/lib/jxl/vardct/lf_group.rb +227 -0
  81. data/lib/jxl/vardct/pass_group.rb +153 -0
  82. data/lib/jxl/vardct/quantizer.rb +27 -0
  83. data/lib/jxl/vardct/reconstructor.rb +294 -0
  84. data/lib/jxl/version.rb +5 -0
  85. data/lib/jxl.rb +216 -0
  86. data/tools/bench.rb +20 -0
  87. data/tools/compare_trace.rb +42 -0
  88. data/tools/conformance_report.rb +23 -0
  89. data/tools/fuzz.rb +51 -0
  90. data/tools/gen_afv_basis.rb +12 -0
  91. data/tools/gen_quant_tables.rb +11 -0
  92. data/tools/gen_upsampling_weights.rb +13 -0
  93. data/tools/libjxl_trace.patch +30 -0
  94. metadata +135 -0
@@ -0,0 +1,371 @@
1
+ # frozen_string_literal: true
2
+
3
+ module JXL
4
+ module Modular
5
+ GroupedState = Data.define(:state, :header, :channels, :shapes, :original_shapes, :shifts,
6
+ :first_group_channel, :group_dim, :metadata)
7
+
8
+ module Decoder
9
+ module_function
10
+
11
+ def decode(data, max_pixels: 100_000_000)
12
+ parsed = Codestream::Reader.new(data, max_pixels:).read_headers
13
+ decode_parsed(parsed)
14
+ end
15
+
16
+ def decode_parsed(parsed, render: false, references: [])
17
+ frame = parsed.frame_header
18
+ raise UnsupportedFeatureError, "VarDCT frame decoding" unless frame.encoding == :modular
19
+
20
+ metadata = parsed.basic_info.metadata
21
+ return decode_grouped(parsed, metadata, render:, references:) unless parsed.sections.one?
22
+
23
+ reader = Bit::Reader.new(parsed.sections.first)
24
+ patches = read_patches(reader, frame, metadata, references)
25
+ dequant = VarDCT::Dequant.read_dc(reader)
26
+ frame_width = frame.encoded_width
27
+ frame_height = frame.encoded_height
28
+ max_nodes = [1024 + (frame_width * frame_height / 4), 1 << 22].min
29
+ state = Stream.read_global_state(reader, max_nodes:)
30
+ count = colour_channels(metadata, frame) + metadata.extra_channels.length
31
+ shapes = Array.new(count) { [frame_width, frame_height] }
32
+ channels = Stream.decode(reader, shapes, state:, stream_id: 0, metadata:, max_nodes:).channels
33
+ convert_float!(channels, metadata.bit_depth)
34
+ image = Image.new(width: frame_width, height: frame_height, channels:, metadata:)
35
+ image = render_image(image, frame, dequant, parsed.transform_data) if render
36
+ Features::Patches.apply(image.channels, patches, references, metadata)
37
+ image
38
+ end
39
+
40
+ def decode_grouped(parsed, metadata, render: false, references: [])
41
+ frame = parsed.frame_header
42
+ frame_width = frame.encoded_width
43
+ frame_height = frame.encoded_height
44
+ group_dim = 128 << frame.group_size_shift
45
+ groups_x = Num.ceil_div(frame_width, group_dim)
46
+ groups_y = Num.ceil_div(frame_height, group_dim)
47
+ dc_groups = Num.ceil_div(Num.ceil_div(frame_width, 8), group_dim) *
48
+ Num.ceil_div(Num.ceil_div(frame_height, 8), group_dim)
49
+
50
+ global_reader = Bit::Reader.new(parsed.sections.first)
51
+ patches = read_patches(global_reader, frame, metadata, references)
52
+ dequant = VarDCT::Dequant.read_dc(global_reader)
53
+ max_nodes = [1024 + (frame_width * frame_height / 4), 1 << 22].min
54
+ state = Stream.read_global_state(global_reader, max_nodes:)
55
+ count = colour_channels(metadata, frame) + metadata.extra_channels.length
56
+ shapes = Array.new(count) { [frame_width, frame_height] }
57
+ grouped = read_grouped_stream(global_reader, shapes, state:, metadata:, group_dim:, max_nodes:)
58
+
59
+ dc_width = Num.ceil_div(Num.ceil_div(frame_width, 8), group_dim)
60
+ dc_groups.times do |group|
61
+ decode_grouped_dc!(Bit::Reader.new(parsed.sections.fetch(1 + group)), grouped,
62
+ group_index: group, groups_x: dc_width, stream_id: 1 + dc_groups + group)
63
+ end
64
+ group_start = 2 + dc_groups
65
+ group_count = groups_x * groups_y
66
+ frame.passes.count.times do |pass|
67
+ group_count.times do |group|
68
+ stream_id = 1 + (3 * dc_groups) + 17 + (group_count * pass) + group
69
+ section = parsed.sections.fetch(group_start + (group_count * pass) + group)
70
+ decode_grouped_pass!(Bit::Reader.new(section), grouped, frame.passes,
71
+ pass_index: pass, group_index: group, groups_x:, stream_id:)
72
+ end
73
+ end
74
+ channels = finish_grouped_stream(grouped)
75
+ convert_float!(channels, metadata.bit_depth)
76
+ image = Image.new(width: frame_width, height: frame_height, channels:, metadata:)
77
+ image = render_image(image, frame, dequant, parsed.transform_data) if render
78
+ Features::Patches.apply(image.channels, patches, references, metadata)
79
+ image
80
+ end
81
+ private_class_method :decode_grouped
82
+
83
+ def read_grouped_stream(reader, original_shapes, state:, metadata:, group_dim:,
84
+ shifts: Array.new(original_shapes.length) { [0, 0] }, max_nodes: 1 << 20)
85
+ header = GroupHeader.read(reader)
86
+ shapes, meta_channels, shifts =
87
+ TransformOps.encoded_shapes_for(original_shapes.map(&:dup), header.transforms, shifts: shifts.map(&:dup))
88
+ tree, entropy = Stream.code(reader, header, state, shapes, max_nodes)
89
+ channels = shapes.map { |width, height| Plane.new(width, height) }
90
+ first_group_channel = shapes.each_index.find do |channel|
91
+ channel >= meta_channels && shapes[channel].any? { _1 > group_dim }
92
+ end || shapes.length
93
+ global_indices = (0...first_group_channel).reject { |channel| shapes[channel].include?(0) }
94
+ unless global_indices.empty?
95
+ entropy.start(distance_multiplier: global_indices.map { |channel| shapes[channel][0] }.max)
96
+ global_indices.each do |channel|
97
+ width, height = shapes[channel]
98
+ channels[channel] = decode_channel(entropy, tree, header.weighted, width, height, channel,
99
+ references: channels.take(channel))
100
+ end
101
+ entropy.final_state!
102
+ end
103
+ GroupedState.new(state:, header:, channels:, shapes:, original_shapes:, shifts:,
104
+ first_group_channel:, group_dim:, metadata:)
105
+ end
106
+
107
+ def decode_grouped_dc!(reader, grouped, group_index:, groups_x:, stream_id:)
108
+ size = grouped.group_dim * 8
109
+ decode_grouped_bracket!(reader, grouped, group_index:, groups_x:, size:,
110
+ min_shift: 3, max_shift: 1000, stream_id:)
111
+ end
112
+
113
+ def decode_grouped_pass!(reader, grouped, passes, pass_index:, group_index:, groups_x:, stream_id:)
114
+ min_shift, max_shift = pass_shift_range(passes, pass_index)
115
+ decode_grouped_bracket!(reader, grouped, group_index:, groups_x:, size: grouped.group_dim,
116
+ min_shift:, max_shift:, stream_id:)
117
+ end
118
+
119
+ def finish_grouped_stream(grouped)
120
+ grouped.header.transforms.reverse_each do |transform|
121
+ inverse_transform!(grouped.channels, transform, grouped.header.weighted, grouped.metadata)
122
+ end
123
+ unless grouped.channels.map { [_1.width, _1.height] } == grouped.original_shapes
124
+ raise CorruptError, "modular transforms changed the requested stream shape"
125
+ end
126
+
127
+ grouped.channels
128
+ end
129
+
130
+ def inverse_transforms!(channels, header, metadata)
131
+ header.transforms.reverse_each do |transform|
132
+ inverse_transform!(channels, transform, header.weighted, metadata)
133
+ end
134
+ channels
135
+ end
136
+
137
+ def decode_channels(entropy, tree, weighted, shapes, stream_id)
138
+ shapes.each_with_index.with_object([]) do |((width, height), channel), channels|
139
+ channels << decode_channel(entropy, tree, weighted, width, height, channel, stream_id,
140
+ references: channels)
141
+ end
142
+ end
143
+
144
+ def section_shape(shape, shifts, channel, x, y, size)
145
+ hshift, vshift = shifts
146
+ target_x = x >> hshift
147
+ target_y = y >> vshift
148
+ width = [size >> hshift, shape[0] - target_x].min
149
+ height = [size >> vshift, shape[1] - target_y].min
150
+ return if width <= 0 || height <= 0
151
+
152
+ [channel, target_x, target_y, width, height]
153
+ end
154
+ private_class_method :section_shape
155
+
156
+ def pass_shift_range(passes, target)
157
+ minimum = 3
158
+ maximum = 2
159
+ passes.count.times do |pass|
160
+ passes.last_passes.each_with_index do |last, index|
161
+ minimum = passes.downsamples[index].bit_length - 1 if pass == last
162
+ end
163
+ minimum = 0 if pass == passes.count - 1
164
+ return [minimum, maximum] if pass == target
165
+
166
+ maximum = minimum - 1
167
+ end
168
+ end
169
+ private_class_method :pass_shift_range
170
+
171
+ def decode_grouped_bracket!(reader, grouped, group_index:, groups_x:, size:,
172
+ min_shift:, max_shift:, stream_id:)
173
+ x = (group_index % groups_x) * size
174
+ y = (group_index / groups_x) * size
175
+ selected = (grouped.first_group_channel...grouped.shapes.length).filter_map do |channel|
176
+ shift = grouped.shifts[channel].min
177
+ next unless shift.between?(min_shift, max_shift)
178
+
179
+ section_shape(grouped.shapes[channel], grouped.shifts[channel], channel, x, y, size)
180
+ end
181
+ decode_group_section!(reader, grouped.state, grouped.channels, selected, stream_id, grouped.metadata)
182
+ end
183
+ private_class_method :decode_grouped_bracket!
184
+
185
+ def decode_group_section!(reader, state, channels, selected, stream_id, metadata)
186
+ return if selected.empty?
187
+
188
+ shapes = selected.map { |(_, _, _, width, height)| [width, height] }
189
+ decoded = Stream.decode(reader, shapes, state:, stream_id:, metadata:).channels
190
+ selected.zip(decoded).each do |(channel, x, y, _, _), plane|
191
+ blit!(plane, channels[channel], x, y)
192
+ end
193
+ end
194
+ private_class_method :decode_group_section!
195
+
196
+ def decode_channel(entropy, tree, weighted_header, width, height, channel, group = 0, references: [])
197
+ plane = Plane.new(width, height)
198
+ weighted = Weighted.new(weighted_header, width) if tree.weighted?
199
+ references = [] unless tree.nodes.any? { !_1.leaf? && _1.property >= 16 }
200
+ height.times do |y|
201
+ previous_gradient = 0
202
+ width.times do |x|
203
+ weighted_guess = weighted&.predict(plane, x, y)
204
+ properties = Predictor.properties(plane, x, y,
205
+ channel:, group:, previous_gradient:,
206
+ weighted_property: weighted&.property || 0, references:)
207
+ previous_gradient = properties[9]
208
+ leaf = tree.leaf(properties)
209
+ begin
210
+ residual = Num.unpack_signed(entropy.read_uint(leaf.context))
211
+ rescue Error => e
212
+ raise e.class, "#{e.message} at channel #{channel}, pixel #{x},#{y}"
213
+ end
214
+ prediction = if leaf.predictor.zero?
215
+ 0
216
+ elsif leaf.predictor == 6
217
+ weighted_guess
218
+ else
219
+ Predictor.predict(leaf.predictor, plane, x, y)
220
+ end
221
+ guess = prediction + leaf.offset
222
+ plane[x, y] = Num.i32((residual * leaf.multiplier) + guess)
223
+ weighted&.update(plane[x, y], x)
224
+ end
225
+ end
226
+ plane
227
+ end
228
+ private_class_method :decode_channel
229
+
230
+ def blit!(source, target, x, y)
231
+ source.height.times do |row|
232
+ target.data[((y + row) * target.width) + x, source.width] =
233
+ source.data[row * source.width, source.width]
234
+ end
235
+ end
236
+ private_class_method :blit!
237
+
238
+ def colour_channels(metadata, frame)
239
+ grayscale = metadata.colour_encoding.colour_space == 1
240
+ grayscale && frame.colour_transform == :none ? 1 : 3
241
+ end
242
+ private_class_method :colour_channels
243
+
244
+ def read_dc_quantization(reader)
245
+ return if Bit::Field.read_bool(reader)
246
+
247
+ 3.times do
248
+ value = Bit::Field.f16(reader) / 128.0
249
+ raise CorruptError, "invalid DC quantization" unless value.positive?
250
+ end
251
+ end
252
+ private_class_method :read_dc_quantization
253
+
254
+ def read_patches(reader, frame, metadata, references)
255
+ return unless frame.flags.anybits?(2)
256
+
257
+ Features::Patches.read(reader, frame.encoded_width, frame.encoded_height,
258
+ metadata.extra_channels.length, references)
259
+ end
260
+ private_class_method :read_patches
261
+
262
+ def render_image(image, frame, dequant, transform_data)
263
+ source = image.channels
264
+ colour_count = colour_channels(image.metadata, frame)
265
+ output = if frame.colour_transform == :xyb
266
+ Array.new(3) { Plane.new(image.width, image.height) }.tap do |planes|
267
+ source[0].data.length.times do |index|
268
+ planes[0].data[index] = source[1].data[index] * dequant.dc[0]
269
+ planes[1].data[index] = source[0].data[index] * dequant.dc[1]
270
+ planes[2].data[index] = (source[2].data[index] + source[0].data[index]) * dequant.dc[2]
271
+ end
272
+ end
273
+ else
274
+ normalize_integer(source.first(colour_count), image.metadata.bit_depth)
275
+ end
276
+ extras = source.drop(colour_count).each_with_index.map do |plane, index|
277
+ depth = image.metadata.extra_channels[index].bit_depth
278
+ normalize_integer([plane], depth).first
279
+ end
280
+ filtered = if colour_count == 1
281
+ grey = if frame.loop_filter.gaborish
282
+ Filter::Gaborish.convolve(output.first, frame.loop_filter.gaborish_weights.first)
283
+ else
284
+ output.first
285
+ end
286
+ Array.new(3, grey) + extras
287
+ else
288
+ Filter::Gaborish.apply(output + extras, frame.loop_filter)
289
+ end
290
+ filtered = Filter::EPF.apply_modular(filtered, frame.loop_filter)
291
+ filtered = Features::Upsampling.apply(filtered, frame, transform_data) unless frame.type == 1
292
+ filtered = [filtered.first] + filtered.drop(3) if colour_count == 1
293
+ dimensions = if frame.type == 1
294
+ { width: image.width, height: image.height }
295
+ else
296
+ { width: frame.width, height: frame.height }
297
+ end
298
+ image.with(**dimensions, channels: filtered)
299
+ end
300
+ private_class_method :render_image
301
+
302
+ def normalize_integer(planes, depth)
303
+ return planes if depth.floating_point
304
+
305
+ scale = 1.0 / ((1 << depth.bits_per_sample) - 1)
306
+ planes.map do |plane|
307
+ Plane.new(plane.width, plane.height).tap { _1.data.replace(plane.data.map { |value| value * scale }) }
308
+ end
309
+ end
310
+ private_class_method :normalize_integer
311
+
312
+ def inverse_transform!(channels, transform, weighted, metadata)
313
+ if transform.id.zero?
314
+ inverse_rct!(channels, transform)
315
+ else
316
+ TransformOps.inverse!(channels, transform, weighted, metadata.bit_depth.bits_per_sample)
317
+ end
318
+ end
319
+ private_class_method :inverse_transform!
320
+
321
+ def convert_float!(channels, bit_depth)
322
+ return unless bit_depth.floating_point
323
+ unless bit_depth.bits_per_sample == 32 && bit_depth.exponent_bits == 8
324
+ raise UnsupportedFeatureError, "non-binary32 floating point samples"
325
+ end
326
+
327
+ channels.each do |plane|
328
+ plane.data.map! { |value| [Num.u32(value)].pack("L<").unpack1("e") }
329
+ end
330
+ end
331
+ private_class_method :convert_float!
332
+
333
+ def inverse_rct!(channels, transform)
334
+ first = transform.begin_channel
335
+ raise CorruptError, "RCT channel range" if first + 2 >= channels.length
336
+
337
+ type = transform.rct_type % 7
338
+ permutation = transform.rct_type / 7
339
+ return permute!(channels, first, permutation) if type.zero?
340
+
341
+ targets = [permutation % 3, (permutation + 1 + (permutation / 3)) % 3,
342
+ (permutation + 2 - (permutation / 3)) % 3]
343
+ source = channels.slice(first, 3)
344
+ source.first.data.length.times do |index|
345
+ a, b, c = source.map { |plane| plane.data[index] }
346
+ if type == 6
347
+ tmp = Num.i32(a - (c >> 1))
348
+ blue = Num.i32(tmp - (b >> 1))
349
+ values = [Num.i32(blue + b), Num.i32(tmp + c), blue]
350
+ else
351
+ c = Num.i32(c + a) if type.odd?
352
+ second = type >> 1
353
+ b = Num.i32(b + a) if second == 1
354
+ b = Num.i32(b + (Num.i32(a + c) >> 1)) if second == 2
355
+ values = [a, b, c]
356
+ end
357
+ targets.each_with_index { |target, i| channels[first + target].data[index] = values[i] }
358
+ end
359
+ end
360
+ private_class_method :inverse_rct!
361
+
362
+ def permute!(channels, first, permutation)
363
+ source = channels.slice(first, 3)
364
+ targets = [permutation % 3, (permutation + 1 + (permutation / 3)) % 3,
365
+ (permutation + 2 - (permutation / 3)) % 3]
366
+ targets.each_with_index { |target, index| channels[first + target] = source[index] }
367
+ end
368
+ private_class_method :permute!
369
+ end
370
+ end
371
+ end
@@ -0,0 +1,82 @@
1
+ # frozen_string_literal: true
2
+
3
+ module JXL
4
+ module Modular
5
+ WeightedHeader = Data.define(:p, :weights)
6
+ Squeeze = Data.define(:horizontal, :in_place, :begin_channel, :num_channels)
7
+ Transform = Data.define(:id, :begin_channel, :rct_type, :num_channels, :num_colors, :num_deltas,
8
+ :predictor, :squeezes)
9
+ GroupHeader = Data.define(:use_global_tree, :weighted, :transforms) do
10
+ def self.read(reader)
11
+ use_global_tree = Bit::Field.read_bool(reader)
12
+ weighted = read_weighted(reader)
13
+ count = Bit::Field.u32(reader, Bit::Field.val(0), Bit::Field.val(1),
14
+ Bit::Field.bits_offset(4, 2), Bit::Field.bits_offset(8, 18))
15
+ raise ResourceLimitError, "too many modular transforms" if count > 256
16
+
17
+ transforms = Array.new(count) { read_transform(reader) }
18
+ new(use_global_tree:, weighted:, transforms:)
19
+ end
20
+
21
+ def self.read_weighted(reader)
22
+ return WeightedHeader.new(p: [16, 10, 7, 7, 7, 0, 0], weights: [13, 12, 12, 12]) if Bit::Field.read_bool(reader)
23
+
24
+ WeightedHeader.new(p: Array.new(7) { reader.read(5) }, weights: Array.new(4) { reader.read(4) })
25
+ end
26
+ private_class_method :read_weighted
27
+
28
+ def self.read_transform(reader)
29
+ id = Bit::Field.u32(reader, Bit::Field.val(0), Bit::Field.val(1), Bit::Field.val(2), Bit::Field.val(3))
30
+ raise CorruptError, "invalid modular transform" if id == 3
31
+
32
+ return read_squeeze(reader, id) if id == 2
33
+
34
+ begin_channel = read_channel(reader)
35
+ return read_palette(reader, id, begin_channel) if id == 1
36
+
37
+ rct_type = Bit::Field.u32(reader, Bit::Field.val(6), Bit::Field.bits(2),
38
+ Bit::Field.bits_offset(4, 2), Bit::Field.bits_offset(6, 10))
39
+ raise CorruptError, "invalid RCT transform" if rct_type >= 42
40
+
41
+ Transform.new(id:, begin_channel:, rct_type:, num_channels: nil, num_colors: nil, num_deltas: nil,
42
+ predictor: nil, squeezes: [])
43
+ end
44
+ private_class_method :read_transform
45
+
46
+ def self.read_palette(reader, id, begin_channel)
47
+ num_channels = Bit::Field.u32(reader, Bit::Field.val(1), Bit::Field.val(3), Bit::Field.val(4),
48
+ Bit::Field.bits_offset(13, 1))
49
+ num_colors = Bit::Field.u32(reader, Bit::Field.bits_offset(8, 0), Bit::Field.bits_offset(10, 256),
50
+ Bit::Field.bits_offset(12, 1280), Bit::Field.bits_offset(16, 5376))
51
+ num_deltas = Bit::Field.u32(reader, Bit::Field.val(0), Bit::Field.bits_offset(8, 1),
52
+ Bit::Field.bits_offset(10, 257), Bit::Field.bits_offset(16, 1281))
53
+ predictor = reader.read(4)
54
+ raise CorruptError, "invalid palette predictor" if predictor >= 14
55
+
56
+ Transform.new(id:, begin_channel:, rct_type: nil, num_channels:, num_colors:, num_deltas:,
57
+ predictor:, squeezes: [])
58
+ end
59
+ private_class_method :read_palette
60
+
61
+ def self.read_squeeze(reader, id)
62
+ count = Bit::Field.u32(reader, Bit::Field.val(0), Bit::Field.bits_offset(4, 1),
63
+ Bit::Field.bits_offset(6, 9), Bit::Field.bits_offset(8, 41))
64
+ squeezes = Array.new(count) do
65
+ Squeeze.new(horizontal: Bit::Field.read_bool(reader), in_place: Bit::Field.read_bool(reader),
66
+ begin_channel: read_channel(reader),
67
+ num_channels: Bit::Field.u32(reader, Bit::Field.val(1), Bit::Field.val(2),
68
+ Bit::Field.val(3), Bit::Field.bits_offset(4, 4)))
69
+ end
70
+ Transform.new(id:, begin_channel: nil, rct_type: nil, num_channels: nil, num_colors: nil,
71
+ num_deltas: nil, predictor: nil, squeezes:)
72
+ end
73
+ private_class_method :read_squeeze
74
+
75
+ def self.read_channel(reader)
76
+ Bit::Field.u32(reader, Bit::Field.bits(3), Bit::Field.bits_offset(6, 8),
77
+ Bit::Field.bits_offset(10, 72), Bit::Field.bits_offset(13, 1096))
78
+ end
79
+ private_class_method :read_channel
80
+ end
81
+ end
82
+ end
@@ -0,0 +1,95 @@
1
+ # frozen_string_literal: true
2
+
3
+ module JXL
4
+ module Modular
5
+ MANode = Data.define(:property, :split_value, :left, :right, :predictor, :offset, :multiplier, :context) do
6
+ def leaf? = property == -1
7
+ end
8
+
9
+ class MATree
10
+ attr_reader :nodes, :leaf_count
11
+
12
+ def initialize(nodes, leaf_count)
13
+ @nodes = nodes
14
+ @leaf_count = leaf_count
15
+ end
16
+
17
+ def self.read(reader, max_nodes: 1 << 20)
18
+ decoder = Entropy::Decoder.read(reader, 6)
19
+ nodes = []
20
+ leaves = 0
21
+ pending = 1
22
+ while pending.positive?
23
+ raise ResourceLimitError, "MA tree is too large" if nodes.length >= max_nodes
24
+
25
+ pending -= 1
26
+ property = decoder.read_uint(1) - 1
27
+ raise CorruptError, "invalid MA tree property" unless property.between?(-1, 255)
28
+
29
+ if property == -1
30
+ nodes << read_leaf(decoder, leaves)
31
+ leaves += 1
32
+ else
33
+ split = Num.unpack_signed(decoder.read_uint(0))
34
+ left = nodes.length + pending + 1
35
+ nodes << MANode.new(property:, split_value: split, left:, right: left + 1,
36
+ predictor: nil, offset: nil, multiplier: nil, context: nil)
37
+ pending += 2
38
+ end
39
+ end
40
+ decoder.final_state!
41
+ validate!(nodes)
42
+ new(nodes, leaves)
43
+ end
44
+
45
+ def leaf(properties)
46
+ index = 0
47
+ loop do
48
+ node = nodes.fetch(index)
49
+ return node if node.leaf?
50
+
51
+ index = (properties[node.property] || 0) > node.split_value ? node.left : node.right
52
+ end
53
+ end
54
+
55
+ def inspect_tree
56
+ nodes.each_with_index.map do |node, index|
57
+ if node.leaf?
58
+ "#{index}: leaf ctx=#{node.context} pred=#{node.predictor} offset=#{node.offset} mul=#{node.multiplier}"
59
+ else
60
+ "#{index}: p#{node.property}>#{node.split_value} ? #{node.left} : #{node.right}"
61
+ end
62
+ end.join("\n")
63
+ end
64
+
65
+ def weighted? = nodes.any? { _1.property == 15 || _1.predictor == 6 }
66
+
67
+ def self.read_leaf(decoder, context)
68
+ predictor = decoder.read_uint(2)
69
+ raise CorruptError, "invalid modular predictor" unless predictor.between?(0, 13)
70
+
71
+ offset = Num.unpack_signed(decoder.read_uint(3))
72
+ multiplier_log = decoder.read_uint(4)
73
+ raise CorruptError, "invalid modular multiplier" if multiplier_log >= 31
74
+
75
+ multiplier_bits = decoder.read_uint(5)
76
+ raise CorruptError, "invalid modular multiplier" if multiplier_bits >= (1 << (31 - multiplier_log)) - 1
77
+
78
+ MANode.new(property: -1, split_value: nil, left: nil, right: nil, predictor:, offset:,
79
+ multiplier: (multiplier_bits + 1) << multiplier_log, context:)
80
+ end
81
+ private_class_method :read_leaf
82
+
83
+ def self.validate!(nodes)
84
+ nodes.each do |node|
85
+ next if node.leaf?
86
+
87
+ unless node.left.between?(0, nodes.length - 1) && node.right.between?(0, nodes.length - 1)
88
+ raise CorruptError, "invalid MA tree child"
89
+ end
90
+ end
91
+ end
92
+ private_class_method :validate!
93
+ end
94
+ end
95
+ end
@@ -0,0 +1,107 @@
1
+ # frozen_string_literal: true
2
+
3
+ module JXL
4
+ module Modular
5
+ module Predictor
6
+ module_function
7
+
8
+ def properties(plane, x, y, channel:, group: 0, previous_gradient: 0, weighted_property: 0, references: [])
9
+ data = plane.data
10
+ width = plane.width
11
+ index = (y * width) + x
12
+ left = if x.positive?
13
+ data[index - 1]
14
+ elsif y.positive?
15
+ data[index - width]
16
+ else
17
+ 0
18
+ end
19
+ top = y.positive? ? data[index - width] : left
20
+ topleft = x.positive? && y.positive? ? data[index - width - 1] : left
21
+ topright = y.positive? && x + 1 < width ? data[index - width + 1] : top
22
+ leftleft = x > 1 ? data[index - 2] : left
23
+ toptop = y > 1 ? data[index - (2 * width)] : top
24
+ values = [channel, group, y, x, top.abs, left.abs, top, left,
25
+ left - previous_gradient, left + top - topleft,
26
+ left - topleft, topleft - top, top - topright,
27
+ top - toptop, left - leftleft, weighted_property]
28
+ references.reverse_each do |reference|
29
+ next unless reference.width == plane.width && reference.height == plane.height
30
+
31
+ value = reference[x, y]
32
+ ref_left = x.positive? ? reference[x - 1, y] : 0
33
+ ref_top = y.positive? ? reference[x, y - 1] : ref_left
34
+ ref_topleft = x.positive? && y.positive? ? reference[x - 1, y - 1] : ref_left
35
+ residual = value - gradient(ref_left, ref_top, ref_topleft)
36
+ values.push(value.abs, value, residual.abs, residual)
37
+ end
38
+ values.map { Num.i32(_1) }
39
+ end
40
+
41
+ def predict(kind, plane, x, y)
42
+ data = plane.data
43
+ width = plane.width
44
+ index = (y * width) + x
45
+ left = if x.positive?
46
+ data[index - 1]
47
+ elsif y.positive?
48
+ data[index - width]
49
+ else
50
+ 0
51
+ end
52
+ top = y.positive? ? data[index - width] : left
53
+ topleft = x.positive? && y.positive? ? data[index - width - 1] : left
54
+ topright = y.positive? && x + 1 < width ? data[index - width + 1] : top
55
+ case kind
56
+ when 0 then 0
57
+ when 1 then left
58
+ when 2 then top
59
+ when 3 then Num.cdiv(left + top, 2)
60
+ when 4 then select(left, top, topleft)
61
+ when 5 then gradient(left, top, topleft)
62
+ when 7 then topright
63
+ when 8 then topleft
64
+ when 9 then x > 1 ? data[index - 2] : left
65
+ when 10 then Num.cdiv(left + topleft, 2)
66
+ when 11 then Num.cdiv(topleft + top, 2)
67
+ when 12 then Num.cdiv(top + topright, 2)
68
+ when 13
69
+ toptop = y > 1 ? data[index - (2 * width)] : top
70
+ leftleft = x > 1 ? data[index - 2] : left
71
+ toprightright = y.positive? && x + 2 < width ? data[index - width + 2] : topright
72
+ Num.cdiv((6 * top) - (2 * toptop) + (7 * left) + leftleft + toprightright + (3 * topright) + 8, 16)
73
+ else
74
+ raise UnsupportedFeatureError, "weighted modular predictor"
75
+ end
76
+ end
77
+
78
+ def gradient(left, top, topleft) = (left + top - topleft).clamp([left, top].min, [left, top].max)
79
+
80
+ def select(left, top, topleft)
81
+ predicted = left + top - topleft
82
+ (predicted - left).abs < (predicted - top).abs ? left : top
83
+ end
84
+
85
+ def neighbors(plane, x, y)
86
+ left =
87
+ if x.positive?
88
+ plane[x - 1, y]
89
+ elsif y.positive?
90
+ plane[x, y - 1]
91
+ else
92
+ 0
93
+ end
94
+ top = y.positive? ? plane[x, y - 1] : left
95
+ topright = y.positive? && x + 1 < plane.width ? plane[x + 1, y - 1] : top
96
+ {
97
+ left:, top:,
98
+ topleft: x.positive? && y.positive? ? plane[x - 1, y - 1] : left,
99
+ topright:,
100
+ leftleft: x > 1 ? plane[x - 2, y] : left,
101
+ toptop: y > 1 ? plane[x, y - 2] : top,
102
+ toprightright: y.positive? && x + 2 < plane.width ? plane[x + 2, y - 1] : topright
103
+ }
104
+ end
105
+ end
106
+ end
107
+ end