lib/simd/src/print.zig
daab053ee43316e1809a84551d573ddd1e5bf3d2
1 const std = @import("std");
2 const bfloat = @import("bfloat.zig");
3
4 pub const Window = struct {
5 lane: usize = 0,
6 max_lanes: usize = 7,
7 };
8
9 const Bounds = struct {
10 begin: usize,
11 end: usize,
12 };
13
14 const Big = struct {
15 limbs: [40]u32 = @splat(0),
16 len: usize = 0,
17
18 fn init(value: u128) Big {
19 var result = Big{};
20 var remaining = value;
21 while (remaining != 0) {
22 result.limbs[result.len] = @truncate(remaining);
23 result.len += 1;
24 remaining >>= 32;
25 }
26 return result;
27 }
28
29 fn multiplySmall(self: *Big, factor: u32) void {
30 var carry: u64 = 0;
31 for (self.limbs[0..self.len]) |*limb| {
32 const product = @as(u64, limb.*) * factor + carry;
33 limb.* = @truncate(product);
34 carry = product >> 32;
35 }
36 if (carry != 0) {
37 std.debug.assert(self.len < self.limbs.len);
38 self.limbs[self.len] = @truncate(carry);
39 self.len += 1;
40 }
41 }
42
43 fn multiplyPower5(self: *Big, count: usize) void {
44 for (0..count) |_| self.multiplySmall(5);
45 }
46
47 fn shiftLeft(self: *Big, amount: usize) void {
48 if (self.len == 0 or amount == 0) return;
49 const words = amount / 32;
50 const bits: u5 = @intCast(amount % 32);
51 std.debug.assert(self.len + words + @intFromBool(bits != 0) <= self.limbs.len);
52 var index = self.len;
53 while (index != 0) {
54 index -= 1;
55 self.limbs[index + words] = self.limbs[index];
56 }
57 @memset(self.limbs[0..words], 0);
58 self.len += words;
59 if (bits == 0) return;
60 var carry: u64 = 0;
61 for (self.limbs[0..self.len]) |*limb| {
62 const shifted = (@as(u64, limb.*) << bits) | carry;
63 limb.* = @truncate(shifted);
64 carry = shifted >> 32;
65 }
66 if (carry != 0) {
67 self.limbs[self.len] = @truncate(carry);
68 self.len += 1;
69 }
70 }
71
72 fn shiftedU128(self: *const Big, shift: usize) u128 {
73 var result: u128 = 0;
74 for (0..4) |word| {
75 result |= @as(u128, self.wordAt(shift + word * 32)) << @intCast(word * 32);
76 }
77 return result;
78 }
79
80 fn roundedShift(self: *const Big, shift: usize) u128 {
81 if (shift == 0) return self.shiftedU128(0);
82 var quotient = self.shiftedU128(shift);
83 const half = self.bitAt(shift - 1);
84 if (half and (self.anyBelow(shift - 1) or quotient & 1 != 0)) quotient += 1;
85 return quotient;
86 }
87
88 fn wordAt(self: *const Big, bit: usize) u32 {
89 const word = bit / 32;
90 if (word >= self.len) return 0;
91 const offset: u5 = @intCast(bit % 32);
92 if (offset == 0) return self.limbs[word];
93 var result = self.limbs[word] >> offset;
94 if (word + 1 < self.len) {
95 const remaining: u5 = @intCast(32 - @as(u6, offset));
96 result |= self.limbs[word + 1] << remaining;
97 }
98 return result;
99 }
100
101 fn bitAt(self: *const Big, bit: usize) bool {
102 const word = bit / 32;
103 if (word >= self.len) return false;
104 const offset: u5 = @intCast(bit % 32);
105 return self.limbs[word] & (@as(u32, 1) << offset) != 0;
106 }
107
108 fn anyBelow(self: *const Big, bit: usize) bool {
109 const words = @min(bit / 32, self.len);
110 for (self.limbs[0..words]) |limb| if (limb != 0) return true;
111 const remaining: u5 = @intCast(bit % 32);
112 if (remaining == 0 or words >= self.len) return false;
113 const mask = (@as(u32, 1) << remaining) - 1;
114 return self.limbs[words] & mask != 0;
115 }
116
117 fn divideSmall(self: *Big, divisor: u32) u32 {
118 var remainder: u64 = 0;
119 var index = self.len;
120 while (index != 0) {
121 index -= 1;
122 const value = (remainder << 32) | self.limbs[index];
123 self.limbs[index] = @intCast(value / divisor);
124 remainder = value % divisor;
125 }
126 while (self.len != 0 and self.limbs[self.len - 1] == 0) self.len -= 1;
127 return @intCast(remainder);
128 }
129 };
130
131 const FloatParts = struct {
132 negative: bool,
133 mantissa: u64,
134 exponent: i32,
135 special: enum { finite, infinity, nan },
136 };
137
138 pub fn writeTypeName(
139 writer: *std.Io.Writer,
140 comptime T: type,
141 lane_count: usize,
142 ) std.Io.Writer.Error!void {
143 validateType(T);
144 try writer.print("{c}{d}", .{ typePrefix(T), typeBits(T) });
145 if (lane_count != 1) try writer.print("x{d}", .{lane_count});
146 }
147
148 pub fn writeValue(
149 writer: *std.Io.Writer,
150 value: anytype,
151 ) std.Io.Writer.Error!void {
152 const T = @TypeOf(value);
153 validateType(T);
154 if (T == bfloat.BFloat16) {
155 return writeFloat(writer, @floatCast(value.toF32()), 3, 1e-3);
156 }
157 switch (@typeInfo(T)) {
158 .int => |info| switch (info.bits) {
159 8 => if (info.signedness == .signed)
160 try writer.print("{d}", .{value})
161 else
162 try writer.print("0x{X:0>2}", .{value}),
163 16 => try writer.print("0x{X:0>4}", .{@as(u16, @bitCast(value))}),
164 32 => try writer.print("{d}", .{value}),
165 64 => try writer.print("0x{x:0>16}", .{@as(u64, @bitCast(value))}),
166 128 => {
167 const bits: u128 = @bitCast(value);
168 try writer.print(
169 "0x{x:0>16}_{x:0>16}",
170 .{ @as(u64, @truncate(bits >> 64)), @as(u64, @truncate(bits)) },
171 );
172 },
173 else => unreachable,
174 },
175 .float => |info| switch (info.bits) {
176 16 => try writeFloat(writer, @floatCast(value), 4, 1e-4),
177 32 => try writeFloat(
178 writer,
179 @floatCast(value),
180 9,
181 @as(f64, @floatCast(@as(f32, 1e-6))),
182 ),
183 64 => try writeFloat(writer, value, 18, 1e-9),
184 else => unreachable,
185 },
186 else => unreachable,
187 }
188 }
189
190 pub fn writeArray(
191 writer: *std.Io.Writer,
192 caption: []const u8,
193 values: anytype,
194 options: Window,
195 ) std.Io.Writer.Error!void {
196 const Slice = @TypeOf(values);
197 const pointer = switch (@typeInfo(Slice)) {
198 .pointer => |info| info,
199 else => @compileError("diagnostic arrays require a slice or array pointer"),
200 };
201 if (pointer.size != .slice and pointer.size != .one) {
202 @compileError("diagnostic arrays require a slice or array pointer");
203 }
204 const T = switch (@typeInfo(pointer.child)) {
205 .array => |info| info.child,
206 else => pointer.child,
207 };
208 validateType(T);
209 const slice: []const T = values;
210 const bounds = try writeHeader(writer, T, caption, slice.len, options);
211 for (slice[bounds.begin..bounds.end]) |value| {
212 try writeValue(writer, value);
213 try writer.writeByte(',');
214 }
215 try writeFooter(writer, bounds);
216 }
217
218 pub fn writeVector(
219 writer: *std.Io.Writer,
220 comptime D: type,
221 caption: []const u8,
222 value: D.Vector,
223 options: Window,
224 ) std.Io.Writer.Error!void {
225 const is_bfloat = @hasDecl(D, "is_bfloat16") and D.is_bfloat16;
226 const T = if (is_bfloat) bfloat.BFloat16 else D.Lane;
227 validateType(T);
228 const bounds = try writeHeader(writer, T, caption, D.lane_count, options);
229 const lanes: [D.lane_count]D.Lane = @bitCast(value);
230 for (bounds.begin..bounds.end) |index| {
231 if (is_bfloat) {
232 try writeValue(writer, bfloat.BFloat16.fromBits(lanes[index]));
233 } else {
234 try writeValue(writer, lanes[index]);
235 }
236 try writer.writeByte(',');
237 }
238 try writeFooter(writer, bounds);
239 }
240
241 fn writeHeader(
242 writer: *std.Io.Writer,
243 comptime T: type,
244 caption: []const u8,
245 lane_count: usize,
246 options: Window,
247 ) std.Io.Writer.Error!Bounds {
248 const begin = options.lane -| 2;
249 const end = @min(begin +| options.max_lanes, lane_count);
250 try writeTypeName(writer, T, lane_count);
251 try writer.print(" {s} [{d}+ ->]:\n ", .{ caption, begin });
252 return .{ .begin = @min(begin, lane_count), .end = end };
253 }
254
255 fn writeFooter(writer: *std.Io.Writer, bounds: Bounds) std.Io.Writer.Error!void {
256 if (bounds.begin >= bounds.end) try writer.writeAll("(out of bounds)");
257 try writer.writeByte('\n');
258 }
259
260 fn writeFloat(
261 writer: *std.Io.Writer,
262 value: f64,
263 comptime precision: usize,
264 threshold: f64,
265 ) std.Io.Writer.Error!void {
266 var buffer: [400]u8 = undefined;
267 const text = if (@abs(value) < threshold)
268 renderScientific(&buffer, value, precision)
269 else
270 renderFixed(&buffer, value, precision);
271 try writer.writeAll(text[0..@min(text.len, 99)]);
272 }
273
274 fn renderScientific(
275 buffer: []u8,
276 value: f64,
277 comptime precision: usize,
278 ) []const u8 {
279 const parts = decompose(value);
280 if (parts.special != .finite) return renderSpecial(buffer, parts);
281 var exponent: i32 = if (parts.mantissa == 0) 0 else decimalExponent(value);
282 var significand = scientificSignificand(parts, precision, exponent);
283 const lower = power10(precision);
284 const upper = lower * 10;
285 if (parts.mantissa != 0 and significand < lower) {
286 exponent -= 1;
287 significand = scientificSignificand(parts, precision, exponent);
288 }
289 if (significand >= upper) {
290 significand /= 10;
291 exponent += 1;
292 }
293 var writer = std.Io.Writer.fixed(buffer);
294 if (parts.negative) writer.writeByte('-') catch unreachable;
295 var digits: [19]u8 = undefined;
296 var remaining = significand;
297 var index = precision + 1;
298 while (index != 0) {
299 index -= 1;
300 digits[index] = '0' + @as(u8, @intCast(remaining % 10));
301 remaining /= 10;
302 }
303 writer.writeByte(digits[0]) catch unreachable;
304 if (precision != 0) {
305 writer.writeByte('.') catch unreachable;
306 writer.writeAll(digits[1 .. precision + 1]) catch unreachable;
307 }
308 writer.writeByte('E') catch unreachable;
309 if (exponent < 0) {
310 writer.writeByte('-') catch unreachable;
311 } else {
312 writer.writeByte('+') catch unreachable;
313 }
314 writer.print("{d:0>2}", .{@abs(exponent)}) catch unreachable;
315 return writer.buffered();
316 }
317
318 fn renderFixed(
319 buffer: []u8,
320 value: f64,
321 comptime precision: usize,
322 ) []const u8 {
323 const parts = decompose(value);
324 if (parts.special != .finite) return renderSpecial(buffer, parts);
325 var scaled = Big.init(parts.mantissa);
326 scaled.multiplyPower5(precision);
327 const binary_exponent = parts.exponent + @as(i32, @intCast(precision));
328 if (binary_exponent >= 0) {
329 scaled.shiftLeft(@intCast(binary_exponent));
330 } else {
331 scaled = Big.init(scaled.roundedShift(@intCast(-binary_exponent)));
332 }
333 var digits_buffer: [400]u8 = undefined;
334 const digits = renderBigDecimal(&digits_buffer, scaled);
335 var writer = std.Io.Writer.fixed(buffer);
336 if (parts.negative) writer.writeByte('-') catch unreachable;
337 if (precision == 0) {
338 writer.writeAll(digits) catch unreachable;
339 } else if (digits.len > precision) {
340 const point = digits.len - precision;
341 writer.writeAll(digits[0..point]) catch unreachable;
342 writer.writeByte('.') catch unreachable;
343 writer.writeAll(digits[point..]) catch unreachable;
344 } else {
345 writer.writeAll("0.") catch unreachable;
346 for (0..precision - digits.len) |_| writer.writeByte('0') catch unreachable;
347 writer.writeAll(digits) catch unreachable;
348 }
349 return writer.buffered();
350 }
351
352 fn renderSpecial(buffer: []u8, parts: FloatParts) []const u8 {
353 var writer = std.Io.Writer.fixed(buffer);
354 if (parts.negative) writer.writeByte('-') catch unreachable;
355 writer.writeAll(if (parts.special == .nan) "nan" else "inf") catch unreachable;
356 return writer.buffered();
357 }
358
359 fn renderBigDecimal(buffer: []u8, value: Big) []const u8 {
360 var remaining = value;
361 if (remaining.len == 0) {
362 buffer[0] = '0';
363 return buffer[0..1];
364 }
365 var chunks: [40]u32 = undefined;
366 var count: usize = 0;
367 while (remaining.len != 0) {
368 chunks[count] = remaining.divideSmall(1_000_000_000);
369 count += 1;
370 }
371 var writer = std.Io.Writer.fixed(buffer);
372 writer.print("{d}", .{chunks[count - 1]}) catch unreachable;
373 var index = count - 1;
374 while (index != 0) {
375 index -= 1;
376 writer.print("{d:0>9}", .{chunks[index]}) catch unreachable;
377 }
378 return writer.buffered();
379 }
380
381 fn scientificSignificand(
382 parts: FloatParts,
383 comptime precision: usize,
384 decimal_exponent: i32,
385 ) u128 {
386 if (parts.mantissa == 0) return 0;
387 const decimal_shift = @as(i32, @intCast(precision)) - decimal_exponent;
388 std.debug.assert(decimal_shift >= 0);
389 var scaled = Big.init(parts.mantissa);
390 scaled.multiplyPower5(@intCast(decimal_shift));
391 const binary_exponent = parts.exponent + decimal_shift;
392 if (binary_exponent >= 0) {
393 scaled.shiftLeft(@intCast(binary_exponent));
394 return scaled.shiftedU128(0);
395 }
396 return scaled.roundedShift(@intCast(-binary_exponent));
397 }
398
399 fn decimalExponent(value: f64) i32 {
400 var normalized: f128 = @floatCast(@abs(value));
401 var exponent: i32 = 0;
402 while (normalized >= 10) {
403 normalized /= 10;
404 exponent += 1;
405 }
406 while (normalized < 1) {
407 normalized *= 10;
408 exponent -= 1;
409 }
410 return exponent;
411 }
412
413 fn decompose(value: f64) FloatParts {
414 const bits: u64 = @bitCast(value);
415 const exponent_bits: u11 = @truncate(bits >> 52);
416 const fraction = bits & 0x000f_ffff_ffff_ffff;
417 const negative = bits >> 63 != 0;
418 if (exponent_bits == 0x7ff) {
419 return .{
420 .negative = negative,
421 .mantissa = fraction,
422 .exponent = 0,
423 .special = if (fraction == 0) .infinity else .nan,
424 };
425 }
426 if (exponent_bits == 0) {
427 return .{
428 .negative = negative,
429 .mantissa = fraction,
430 .exponent = -1074,
431 .special = .finite,
432 };
433 }
434 return .{
435 .negative = negative,
436 .mantissa = fraction | (@as(u64, 1) << 52),
437 .exponent = @as(i32, exponent_bits) - 1023 - 52,
438 .special = .finite,
439 };
440 }
441
442 fn power10(comptime exponent: usize) u128 {
443 var result: u128 = 1;
444 for (0..exponent) |_| result *= 10;
445 return result;
446 }
447
448 fn typePrefix(comptime T: type) u8 {
449 if (T == bfloat.BFloat16) return 'i';
450 return switch (@typeInfo(T)) {
451 .float => 'f',
452 .int => |info| if (info.signedness == .signed) 'i' else 'u',
453 else => unreachable,
454 };
455 }
456
457 fn typeBits(comptime T: type) usize {
458 return if (T == bfloat.BFloat16) 16 else @bitSizeOf(T);
459 }
460
461 fn validateType(comptime T: type) void {
462 if (T == bfloat.BFloat16) return;
463 switch (@typeInfo(T)) {
464 .int => |info| {
465 if (info.bits != 8 and info.bits != 16 and info.bits != 32 and
466 info.bits != 64 and info.bits != 128)
467 {
468 @compileError("diagnostic integers require 8/16/32/64/128 bits");
469 }
470 if (info.bits == 128 and info.signedness == .signed) {
471 @compileError("Highway diagnostics only support unsigned 128-bit lanes");
472 }
473 },
474 .float => |info| if (info.bits != 16 and info.bits != 32 and info.bits != 64) {
475 @compileError("diagnostic floats require 16/32/64 bits");
476 },
477 else => @compileError("unsupported Highway diagnostic lane type"),
478 }
479 }
480
481 test "pinned Highway diagnostic type and value oracle matches byte-for-byte" {
482 var buffer: [2048]u8 = undefined;
483 var writer = std.Io.Writer.fixed(&buffer);
484 const cases = .{
485 .{ "u8", @as(u8, 0xaf), @as(usize, 8) },
486 .{ "i8", @as(i8, -123), @as(usize, 1) },
487 .{ "u16", @as(u16, 0xabcd), @as(usize, 4) },
488 .{ "i16", @as(i16, -2), @as(usize, 4) },
489 .{ "f16-small", @as(f16, @bitCast(@as(u16, 1))), @as(usize, 4) },
490 .{ "f16-fixed", @as(f16, 1.5), @as(usize, 4) },
491 .{ "bf16-small", bfloat.BFloat16.fromBits(1), @as(usize, 4) },
492 .{ "bf16-fixed", bfloat.BFloat16.fromF32(-2.25), @as(usize, 4) },
493 .{ "u32", @as(u32, std.math.maxInt(u32)), @as(usize, 8) },
494 .{ "i32", @as(i32, std.math.minInt(i32)), @as(usize, 8) },
495 .{ "f32-small", @as(f32, 0.0000005), @as(usize, 8) },
496 .{ "f32-fixed", @as(f32, -12.25), @as(usize, 8) },
497 .{ "f32-tenth", @as(f32, 0.1), @as(usize, 8) },
498 .{ "u64", @as(u64, 0x0123_4567_89ab_cdef), @as(usize, 2) },
499 .{ "i64", @as(i64, -2), @as(usize, 2) },
500 .{ "f64-small", @as(f64, 0.0000000005), @as(usize, 2) },
501 .{ "f64-fixed", @as(f64, -12.25), @as(usize, 2) },
502 .{ "f64-tenth", @as(f64, 0.1), @as(usize, 2) },
503 .{ "f64-max", @as(f64, 0x1.fffffffffffffp1023), @as(usize, 2) },
504 .{ "u128", @as(u128, 0xfedc_ba98_7654_3210_0123_4567_89ab_cdef), @as(usize, 2) },
505 };
506 inline for (cases) |case| {
507 try writer.print("{s} type=", .{case[0]});
508 try writeTypeName(&writer, @TypeOf(case[1]), case[2]);
509 try writer.writeAll(" value=");
510 try writeValue(&writer, case[1]);
511 try writer.writeByte('\n');
512 }
513 try std.testing.expectEqualStrings(
514 \\u8 type=u8x8 value=0xAF
515 \\i8 type=i8 value=-123
516 \\u16 type=u16x4 value=0xABCD
517 \\i16 type=i16x4 value=0xFFFE
518 \\f16-small type=f16x4 value=5.9605E-08
519 \\f16-fixed type=f16x4 value=1.5000
520 \\bf16-small type=i16x4 value=9.184E-41
521 \\bf16-fixed type=i16x4 value=-2.250
522 \\u32 type=u32x8 value=4294967295
523 \\i32 type=i32x8 value=-2147483648
524 \\f32-small type=f32x8 value=4.999999987E-07
525 \\f32-fixed type=f32x8 value=-12.250000000
526 \\f32-tenth type=f32x8 value=0.100000001
527 \\u64 type=u64x2 value=0x0123456789abcdef
528 \\i64 type=i64x2 value=0xfffffffffffffffe
529 \\f64-small type=f64x2 value=5.000000000000000311E-10
530 \\f64-fixed type=f64x2 value=-12.250000000000000000
531 \\f64-tenth type=f64x2 value=0.100000000000000006
532 \\f64-max type=f64x2 value=179769313486231570814527423731704356798070567525844996598917476803157260780028538760589558632766878
533 \\u128 type=u128x2 value=0xfedcba9876543210_0123456789abcdef
534 \\
535 ,
536 writer.buffered(),
537 );
538 }
539
540 test "pinned Highway diagnostic lane windows match byte-for-byte" {
541 const values = [_]u32{ 10, 11, 12, 13, 14, 15, 16, 17 };
542 var buffer: [512]u8 = undefined;
543 var writer = std.Io.Writer.fixed(&buffer);
544 try writeArray(&writer, "lanes", &values, .{ .lane = 4, .max_lanes = 5 });
545 try writeArray(&writer, "oob", &values, .{ .lane = 99 });
546 try writeArray(&writer, "zero", &values, .{ .max_lanes = 0 });
547 try std.testing.expectEqualStrings(
548 \\u32x8 lanes [2+ ->]:
549 \\ 12,13,14,15,16,
550 \\u32x8 oob [97+ ->]:
551 \\ (out of bounds)
552 \\u32x8 zero [0+ ->]:
553 \\ (out of bounds)
554 \\
555 ,
556 writer.buffered(),
557 );
558 }
559
560 test "vector diagnostics preserve array order and bfloat meaning" {
561 const tag = @import("tag.zig");
562 const D = tag.FixedTag(i16, 8);
563 const values = [_]i16{ -4, -3, -2, -1, 0, 1, 2, 3 };
564 const vector: D.Vector = values;
565 var array_buffer: [256]u8 = undefined;
566 var array_writer = std.Io.Writer.fixed(&array_buffer);
567 try writeArray(&array_writer, "vector", &values, .{ .lane = 5, .max_lanes = 4 });
568 var vector_buffer: [256]u8 = undefined;
569 var vector_writer = std.Io.Writer.fixed(&vector_buffer);
570 try writeVector(&vector_writer, D, "vector", vector, .{ .lane = 5, .max_lanes = 4 });
571 try std.testing.expectEqualStrings(array_writer.buffered(), vector_writer.buffered());
572
573 const BD = bfloat.Tag(4);
574 const bvalues = [_]bfloat.BFloat16{
575 bfloat.BFloat16.fromF32(1),
576 bfloat.BFloat16.fromF32(-2.25),
577 bfloat.BFloat16.fromBits(1),
578 bfloat.BFloat16.fromF32(4),
579 };
580 const bvector = bfloat.load(BD, &bvalues);
581 var bbuffer: [256]u8 = undefined;
582 var bwriter = std.Io.Writer.fixed(&bbuffer);
583 try writeVector(&bwriter, BD, "bf16", bvector, .{ .max_lanes = 4 });
584 try std.testing.expectEqualStrings(
585 "i16x4 bf16 [0+ ->]:\n 1.000,-2.250,9.184E-41,4.000,\n",
586 bwriter.buffered(),
587 );
588 }
589
590 test "diagnostic boundaries retain scientific thresholds and empty slices" {
591 var buffer: [512]u8 = undefined;
592 var writer = std.Io.Writer.fixed(&buffer);
593 try writeValue(&writer, @as(f32, -0.0));
594 try writer.writeByte(' ');
595 try writeValue(&writer, @as(f32, @bitCast(@as(u32, 0x3586_37bc))));
596 try writer.writeByte(' ');
597 try writeValue(&writer, @as(f32, @bitCast(@as(u32, 0x3586_37bd))));
598 try writer.writeByte(' ');
599 try writeValue(&writer, @as(f32, @bitCast(@as(u32, 0x3586_37be))));
600 try writer.writeByte(' ');
601 try writeValue(&writer, std.math.inf(f64));
602 try writer.writeByte(' ');
603 try writeValue(&writer, std.math.nan(f64));
604 try writer.writeByte('\n');
605 try writeArray(&writer, "empty", @as([]const u8, &.{}), .{});
606 try std.testing.expectEqualStrings(
607 "-0.000000000E+00 9.999998838E-07 0.000001000 0.000001000 inf nan\nu8x0 empty [0+ ->]:\n (out of bounds)\n",
608 writer.buffered(),
609 );
610 }