Skip to content

Commit aa7dddb

Browse files
committed
Fix: Harden reductions, casts, and bindings against degenerate inputs
Several SIMD kernels and Python-binding paths mishandled degenerate shapes and strides, ranging from a hang to out-of-bounds reads and writes. The reduce kernels (skylake/icelake/haswell/sierra/alder/neonbfdot) hit an infinite loop, SIGFPE, or stack overflow when handed stride_bytes == 0; their `!aligned` serial fallbacks now also catch `stride_elements == 0`. The sub-byte cast oracle in cast/serial.h hard-coded its pack/unpack loops to four bytes regardless of count, reading and writing past the buffer on any odd i4/u4/e2m1 length -- every site now bounds the loop by nk_size_divide_round_up_(count, 2) and guards the odd tail. Rank-0 and empty reductions no longer touch a non-existent element: C++ moments()/minmax() indexed stride_bytes(SIZE_MAX) on a 0-D view, the Python rank-0 path fabricated a zero stride, and empty minmax primed its accumulators from element [0]. The Python bindings gained the missing input validation. Dense-metric `out=` buffers are checked for rank and capacity before writing (an undersized out overflowed the heap); parse_tensor now requires exact inner-axis contiguity, rejecting the negative/zero strides (e.g. x[::-1]) that the old signed `> itemsize` check let walk off the buffer; and DLPack import rejects negative extents and a NULL data pointer. The Rust symmetric matrix verbs reject non-contiguous-row (transposed) views, matching the packed and parallel paths. The test suites were deduplicated and extended over the same edges. The Python *_float/*_integer reduction and arithmetic pairs collapse into parametrized tests (removing dead precise_*/baseline_sum helpers and a stale DLPack helper), and new cases cover empty/0-D tensors, NaN/Inf preservation on dense casts, and block-scaled round-trip / idempotence / byte-identity across all seven formats, plus the degenerate inputs above. Net test lines shrink while coverage grows. Comments were stripped of decorative banner separators and ephemeral numbered "Phase"/"Step"/"Option" labels, whose ordering is implied by code order. Also sets the rustfmt line width to 120 and bumps the Rust MSRV to 1.73 for usize::div_ceil.
1 parent 184fd27 commit aa7dddb

47 files changed

Lines changed: 944 additions & 701 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

.cmake-format.py

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,4 @@
1-
# -----------------------------
2-
# Options effecting formatting.
3-
# -----------------------------
1+
# Options affecting formatting.
42
with section("format"):
53
# How wide to allow formatted cmake files
64
line_width = 120
@@ -29,9 +27,7 @@
2927
# one-per-line when word-wrapping exceeds 2 lines)
3028
max_lines_hwrap = 8
3129

32-
# ----------------------------------
3330
# Options affecting comment handling.
34-
# ----------------------------------
3531
with section("markup"):
3632
# Do not reflow comment text
3733
enable_markup = False

Cargo.toml

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@ links = "numkong"
1818
name = "numkong"
1919
readme = "rust/README.md"
2020
repository = "https://github.com/ashvardanian/NumKong"
21-
rust-version = "1.64" # Introduced Core C FFI in stable Rust
21+
rust-version = "1.73" # `usize::div_ceil` (used for sub-byte / block-scaled size math)
2222
version = "7.7.0"
2323

2424
[lib]

include/numkong/attention/sapphireamx.h

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -788,7 +788,7 @@ NK_PUBLIC void nk_attention_bf16_sapphireamx(nk_bf16_t const *q, void const *kv_
788788
// Extract K block: Kᵀ[head_dim, valid_kv] using bulk extraction
789789
nk_attention_extract_k_block_(k_packed, k_block, kv_h, kvb, valid_kv, head_dim, kv_len);
790790

791-
// Phase 1: Compute S = Q × Kᵀ using AVX-512 FMA
791+
// Compute S = Q × Kᵀ using AVX-512 FMA
792792
for (nk_size_t qi = 0; qi < valid_q; qi++) {
793793
for (nk_size_t ki = 0; ki < valid_kv; ki++) {
794794
__m512 sum_v_f32x16 = _mm512_setzero_ps();
@@ -819,7 +819,7 @@ NK_PUBLIC void nk_attention_bf16_sapphireamx(nk_bf16_t const *q, void const *kv_
819819
for (nk_size_t ki = 0; ki < 16; ki++) { scores[qi * 16 + ki] = NK_F32_MIN; }
820820
}
821821

822-
// Phase 2: Online softmax update
822+
// Online softmax update
823823
__m512 old_max_f32x16 = softmax_state.row_max_f32x16;
824824
nk_attention_softmax_update_(&softmax_state, scores, scale, weights);
825825

@@ -829,7 +829,7 @@ NK_PUBLIC void nk_attention_bf16_sapphireamx(nk_bf16_t const *q, void const *kv_
829829
// Extract V block: V[valid_kv, head_dim] using bulk extraction
830830
nk_attention_extract_v_block_(v_packed, v_block, kv_h, kvb, valid_kv, head_dim, kv_len);
831831

832-
// Phase 3: Compute O += P × V using AVX-512 FMA
832+
// Compute O += P × V using AVX-512 FMA
833833
for (nk_size_t qi = 0; qi < valid_q; qi++) {
834834
nk_size_t d = 0;
835835
// Vectorized loop over head_dim
@@ -931,7 +931,7 @@ NK_PUBLIC void nk_attention_bf16_amx_bc32_sapphireamx(nk_bf16_t const *q, void c
931931
for (nk_size_t kvb = 0; kvb < kv_len; kvb += Bc) {
932932
nk_size_t valid_kv = (kvb + Bc <= kv_len) ? Bc : (kv_len - kvb);
933933

934-
// Phase 1: S = Q × Kᵀ using AMX
934+
// S = Q × Kᵀ using AMX
935935
// Need 2 K tiles per block (each K tile has 16 columns)
936936
nk_size_t k_tile_idx0 = kvb / 16; // First K tile
937937
nk_size_t k_tile_idx1 = (kvb + 16) / 16; // Second K tile
@@ -1038,12 +1038,12 @@ NK_PUBLIC void nk_attention_bf16_amx_bc32_sapphireamx(nk_bf16_t const *q, void c
10381038
for (nk_size_t qi = 0; qi < 16; qi++) { _mm512_store_ps(&scores[qi * 32 + 16], neg_inf_f32x16); }
10391039
}
10401040

1041-
// Phase 2: online softmax (fast degree-4 exp)
1041+
// online softmax (fast degree-4 exp)
10421042
__m512 old_max_f32x16 = softmax_state.row_max_f32x16;
10431043
nk_attention_softmax_update_bc32_fast_(&softmax_state, scores, scale, weights);
10441044
nk_attention_rescale_output_(o_acc, head_dim_padded, old_max_f32x16, softmax_state.row_max_f32x16);
10451045

1046-
// Phase 3: O += P × V using AMX
1046+
// O += P × V using AMX
10471047
// Convert P[16, 32] from F32 to BF16 and pack as A-tile
10481048
for (nk_size_t qi = 0; qi < 16; qi++) {
10491049
for (nk_size_t ki = 0; ki < 32; ki += 16) {
@@ -1213,7 +1213,7 @@ NK_PUBLIC void nk_attention_bf16_amx_optimized_sapphireamx(nk_bf16_t const *q, v
12131213
nk_size_t k_tile_idx0 = kvb / 16;
12141214
nk_size_t k_tile_idx1 = (kvb + 16) / 16;
12151215

1216-
// Phase 1: S = Q × Kᵀ using pre-packed Q tiles
1216+
// S = Q × Kᵀ using pre-packed Q tiles
12171217
_tile_zero(0); // Score cols 0:16
12181218
_tile_zero(3); // Score cols 16:32
12191219

@@ -1263,13 +1263,13 @@ NK_PUBLIC void nk_attention_bf16_amx_optimized_sapphireamx(nk_bf16_t const *q, v
12631263
}
12641264
}
12651265

1266-
// Phase 2: online softmax (fast degree-4 exp)
1266+
// online softmax (fast degree-4 exp)
12671267
__m512 old_max_f32x16 = softmax_state.row_max_f32x16;
12681268
nk_attention_softmax_update_bc32_fast_(&softmax_state, &scores[0][0], scale, &weights[0][0]);
12691269
nk_attention_rescale_output_(&o_acc[0][0], head_dim_padded, old_max_f32x16,
12701270
softmax_state.row_max_f32x16);
12711271

1272-
// Phase 3: O += P × V with hoisted P tile load
1272+
// O += P × V with hoisted P tile load
12731273
// Convert F32 weights to BF16 P tile (once per KV block)
12741274
for (nk_size_t qi = 0; qi < 16; qi++) {
12751275
__m512 p0_f32x16 = _mm512_load_ps(&weights[qi][0]);

include/numkong/attention/sme.h

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -63,14 +63,14 @@
6363
*
6464
* @section attention_sme_history Optimization History
6565
*
66-
* Phase 1 (January 2026): Initial implementation using ZA staging transpose for Q×Kᵀ
66+
* January 2026: Initial implementation using ZA staging transpose for Q×Kᵀ
6767
* and element-wise SVE for P×V. Q and K rows were loaded into ZA0/ZA1 horizontally,
6868
* read back vertically to produce interleaved vectors for BFMOPA. The P×V phase used
6969
* scalar `svmla_f32_x` loops over head_dim for each query-key pair. Softmax used
7070
* degree-4 polynomial exp with per-row horizontal max/sum. Performance: ~25-50 GFLOP/s
7171
* on Apple M4 (bf16, 8 heads, query_len=64, kv_len=4096, head_dim=128).
7272
*
73-
* Phase 2 (February 2026): BFMOPA/FMOPA P×V with pre-packed V in interleaved format.
73+
* February 2026: BFMOPA/FMOPA P×V with pre-packed V in interleaved format.
7474
* Key changes integrated:
7575
* - Q pre-transposed once into a buffer, eliminating per-block ZA staging for Q
7676
* - K pre-packed in interleaved format, enabling pure memory-to-BFMOPA Q×Kᵀ
@@ -1260,7 +1260,7 @@ __arm_new("za") static void nk_attention_f16_sme_streaming_( //
12601260
svread_ver_za32_f32_m(svdup_f32(0), predicate_all_b32x, 0, step));
12611261
}
12621262

1263-
// === Bc=32 main loop (prefill only, skipped for decode) ===
1263+
// Bc=32 main loop, used for prefill only and skipped for decode.
12641264
if (valid_query_count > 1) {
12651265
for (; kv_start + 32 <= kv_len; kv_start += 32, kv_block_index += 2) {
12661266
// Q×K^T: pure memory→FMOPA, no ZA staging for Q or K
@@ -1691,7 +1691,7 @@ __arm_new("za") static void nk_attention_f16_sme_streaming_( //
16911691
}
16921692
}
16931693

1694-
// === Bc=16 tail loop (handles remaining KV positions and decode path) ===
1694+
// Bc=16 tail loop handles the remaining KV positions and the decode path.
16951695
for (; kv_start < kv_len; kv_start += 16, kv_block_index++) {
16961696
nk_size_t const valid_kv = ((kv_start + 16) <= kv_len) ? 16 : (kv_len - kv_start);
16971697

include/numkong/cast/serial.h

Lines changed: 22 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -2012,7 +2012,7 @@ NK_INTERNAL void nk_scalar_buffers_to_f64c_( //
20122012
case nk_i4_k: {
20132013
nk_i4x2_t const *pairs = (nk_i4x2_t const *)from_ptr;
20142014
nk_i8_t unpacked[2];
2015-
for (i = 0; i < 4; ++i) {
2015+
for (i = 0; i < nk_size_divide_round_up_(from_count, 2); ++i) {
20162016
nk_i4x2_to_i8x2_serial(&pairs[i], unpacked);
20172017
to_buffers[i * 2].f64c.real = unpacked[0], to_buffers[i * 2].f64c.imag = 0;
20182018
to_buffers[i * 2 + 1].f64c.real = unpacked[1], to_buffers[i * 2 + 1].f64c.imag = 0;
@@ -2022,7 +2022,7 @@ NK_INTERNAL void nk_scalar_buffers_to_f64c_( //
20222022
case nk_u4_k: {
20232023
nk_u4x2_t const *pairs = (nk_u4x2_t const *)from_ptr;
20242024
nk_u8_t unpacked[2];
2025-
for (i = 0; i < 4; ++i) {
2025+
for (i = 0; i < nk_size_divide_round_up_(from_count, 2); ++i) {
20262026
nk_u4x2_to_u8x2_serial(&pairs[i], unpacked);
20272027
to_buffers[i * 2].f64c.real = unpacked[0], to_buffers[i * 2].f64c.imag = 0;
20282028
to_buffers[i * 2 + 1].f64c.real = unpacked[1], to_buffers[i * 2 + 1].f64c.imag = 0;
@@ -2032,7 +2032,7 @@ NK_INTERNAL void nk_scalar_buffers_to_f64c_( //
20322032
case nk_e2m1_k: {
20332033
nk_e2m1x2_t const *pairs = (nk_e2m1x2_t const *)from_ptr;
20342034
nk_f32_t unpacked[2];
2035-
for (i = 0; i < 4; ++i) {
2035+
for (i = 0; i < nk_size_divide_round_up_(from_count, 2); ++i) {
20362036
nk_e2m1x2_to_f32x2_serial(&pairs[i], unpacked);
20372037
to_buffers[i * 2].f64c.real = (nk_f64_t)unpacked[0], to_buffers[i * 2].f64c.imag = 0;
20382038
to_buffers[i * 2 + 1].f64c.real = (nk_f64_t)unpacked[1], to_buffers[i * 2 + 1].f64c.imag = 0;
@@ -2149,8 +2149,9 @@ NK_INTERNAL void nk_scalar_buffers_from_f64c_( //
21492149
// Sub-byte: i4 - 8 nibbles to 4 bytes, high nibble = even index
21502150
case nk_i4_k: {
21512151
nk_u8_t *p = (nk_u8_t *)to_ptr;
2152-
for (i = 0; i < 4; ++i) {
2153-
nk_f64_t high = from_buffers[i * 2].f64c.real, low = from_buffers[i * 2 + 1].f64c.real;
2152+
for (i = 0; i < nk_size_divide_round_up_(to_count, 2); ++i) {
2153+
nk_f64_t high = from_buffers[i * 2].f64c.real;
2154+
nk_f64_t low = (i * 2 + 1 < to_count) ? from_buffers[i * 2 + 1].f64c.real : 0.0;
21542155
high = high > 7 ? 7 : (high < -8 ? -8 : high);
21552156
low = low > 7 ? 7 : (low < -8 ? -8 : low);
21562157
p[i] = (nk_u8_t)((((nk_i8_t)high & 0x0F) << 4) | ((nk_i8_t)low & 0x0F));
@@ -2159,8 +2160,9 @@ NK_INTERNAL void nk_scalar_buffers_from_f64c_( //
21592160
// Sub-byte: u4 - 8 nibbles to 4 bytes, high nibble = even index
21602161
case nk_u4_k: {
21612162
nk_u8_t *p = (nk_u8_t *)to_ptr;
2162-
for (i = 0; i < 4; ++i) {
2163-
nk_f64_t high = from_buffers[i * 2].f64c.real, low = from_buffers[i * 2 + 1].f64c.real;
2163+
for (i = 0; i < nk_size_divide_round_up_(to_count, 2); ++i) {
2164+
nk_f64_t high = from_buffers[i * 2].f64c.real;
2165+
nk_f64_t low = (i * 2 + 1 < to_count) ? from_buffers[i * 2 + 1].f64c.real : 0.0;
21642166
high = high > 15 ? 15 : (high < 0 ? 0 : high);
21652167
low = low > 15 ? 15 : (low < 0 ? 0 : low);
21662168
p[i] = (nk_u8_t)(((nk_u8_t)high << 4) | (nk_u8_t)low);
@@ -2169,10 +2171,10 @@ NK_INTERNAL void nk_scalar_buffers_from_f64c_( //
21692171
// Sub-byte: e2m1 - 8 nibbles to 4 bytes, high nibble = even index
21702172
case nk_e2m1_k: {
21712173
nk_e2m1x2_t *pairs = (nk_e2m1x2_t *)to_ptr;
2172-
for (i = 0; i < 4; ++i) {
2174+
for (i = 0; i < nk_size_divide_round_up_(to_count, 2); ++i) {
21732175
nk_f32_t paired[2];
21742176
paired[0] = (nk_f32_t)from_buffers[i * 2].f64c.real;
2175-
paired[1] = (nk_f32_t)from_buffers[i * 2 + 1].f64c.real;
2177+
paired[1] = (i * 2 + 1 < to_count) ? (nk_f32_t)from_buffers[i * 2 + 1].f64c.real : 0.0f;
21762178
nk_f32x2_to_e2m1x2_serial(paired, &pairs[i]);
21772179
}
21782180
} break;
@@ -2215,7 +2217,7 @@ NK_INTERNAL void nk_scalar_buffers_to_i64_( //
22152217
// Sub-byte: i4 - 4 bytes to 8 nibbles, sign-extend each nibble
22162218
case nk_i4_k: {
22172219
nk_i4x2_t const *pairs = (nk_i4x2_t const *)from_ptr;
2218-
for (i = 0; i < 4; ++i) {
2220+
for (i = 0; i < nk_size_divide_round_up_(from_count, 2); ++i) {
22192221
nk_i8_t unpacked[2];
22202222
nk_i4x2_to_i8x2_serial(&pairs[i], unpacked);
22212223
to_buffers[i * 2].i64 = unpacked[0];
@@ -2240,7 +2242,7 @@ NK_INTERNAL void nk_scalar_buffers_to_i64_( //
22402242
} break;
22412243
case nk_u4_k: {
22422244
nk_u8_t const *p = (nk_u8_t const *)from_ptr;
2243-
for (i = 0; i < 4; ++i) {
2245+
for (i = 0; i < nk_size_divide_round_up_(from_count, 2); ++i) {
22442246
to_buffers[i * 2].i64 = (nk_i64_t)(p[i] >> 4);
22452247
to_buffers[i * 2 + 1].i64 = (nk_i64_t)(p[i] & 0xF);
22462248
}
@@ -2294,8 +2296,9 @@ NK_INTERNAL void nk_scalar_buffers_from_i64_( //
22942296
// Sub-byte: i4 - 8 nibbles to 4 bytes, clamp [-8,7]
22952297
case nk_i4_k: {
22962298
nk_i4x2_t *p = (nk_i4x2_t *)to_ptr;
2297-
for (i = 0; i < 4; ++i) {
2298-
nk_i64_t high = from_buffers[i * 2].i64, low = from_buffers[i * 2 + 1].i64;
2299+
for (i = 0; i < nk_size_divide_round_up_(to_count, 2); ++i) {
2300+
nk_i64_t high = from_buffers[i * 2].i64;
2301+
nk_i64_t low = (i * 2 + 1 < to_count) ? from_buffers[i * 2 + 1].i64 : 0;
22992302
high = high > 7 ? 7 : (high < -8 ? -8 : high);
23002303
low = low > 7 ? 7 : (low < -8 ? -8 : low);
23012304
p[i] = (nk_u8_t)(((high & 0xF) << 4) | (low & 0xF));
@@ -2332,7 +2335,7 @@ NK_INTERNAL void nk_scalar_buffers_to_u64_( //
23322335
// Sub-byte: u4 - 4 bytes to 8 nibbles, zero-extend
23332336
case nk_u4_k: {
23342337
nk_u4x2_t const *pairs = (nk_u4x2_t const *)from_ptr;
2335-
for (i = 0; i < 4; ++i) {
2338+
for (i = 0; i < nk_size_divide_round_up_(from_count, 2); ++i) {
23362339
nk_u8_t unpacked[2];
23372340
nk_u4x2_to_u8x2_serial(&pairs[i], unpacked);
23382341
to_buffers[i * 2].u64 = unpacked[0];
@@ -2393,8 +2396,9 @@ NK_INTERNAL void nk_scalar_buffers_from_u64_( //
23932396
// Sub-byte: u4 - 8 nibbles to 4 bytes, clamp [0,15]
23942397
case nk_u4_k: {
23952398
nk_u4x2_t *p = (nk_u4x2_t *)to_ptr;
2396-
for (i = 0; i < 4; ++i) {
2397-
nk_u64_t high = from_buffers[i * 2].u64, low = from_buffers[i * 2 + 1].u64;
2399+
for (i = 0; i < nk_size_divide_round_up_(to_count, 2); ++i) {
2400+
nk_u64_t high = from_buffers[i * 2].u64;
2401+
nk_u64_t low = (i * 2 + 1 < to_count) ? from_buffers[i * 2 + 1].u64 : 0;
23982402
high = high > 15 ? 15 : high;
23992403
low = low > 15 ? 15 : low;
24002404
p[i] = (nk_u8_t)((high << 4) | low);
@@ -2861,7 +2865,7 @@ NK_PUBLIC void nk_cast_block_scaled_serial(
28612865
for (nk_size_t chunk_start = 0; chunk_start < count; chunk_start += chunk) {
28622866
nk_size_t chunk_count = (chunk_start + chunk <= count) ? chunk : (count - chunk_start);
28632867

2864-
// --- Decode: source chunk scratch[0..chunk_count) as f32 ---
2868+
// Decode the source chunk into scratch[0..chunk_count) as f32.
28652869
if (from_plain) {
28662870
void const *src = (nk_u8_t const *)from + (chunk_start * from_bits_per_element / NK_BITS_PER_BYTE);
28672871
nk_cast_serial(src, from_format->element_dtype, chunk_count, scratch, nk_f32_k);
@@ -2880,7 +2884,7 @@ NK_PUBLIC void nk_cast_block_scaled_serial(
28802884
}
28812885
}
28822886

2883-
// --- Encode: scratch[0..chunk_count) as f32 destination chunk ---
2887+
// Encode scratch[0..chunk_count) as f32 into the destination chunk.
28842888
if (to_plain) {
28852889
void *dst = (nk_u8_t *)to + (chunk_start * to_bits_per_element / NK_BITS_PER_BYTE);
28862890
nk_cast_serial(scratch, nk_f32_k, chunk_count, dst, to_format->element_dtype);

include/numkong/dot/haswell.h

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1505,7 +1505,7 @@ NK_INTERNAL void nk_dot_through_i32_finalize_haswell_(
15051505
nk_size_t total_dimensions, nk_b128_vec_t *result) {
15061506
nk_unused_(total_dimensions);
15071507
// ILP-optimized 4-way horizontal reduction for i32 in AVX2
1508-
// Step 1: 8->4 for all 4 states
1508+
// 8->4 for all 4 states
15091509
__m128i sum_a_i32x4 = _mm_add_epi32(_mm256_castsi256_si128(state_a->sum_i32x8),
15101510
_mm256_extracti128_si256(state_a->sum_i32x8, 1));
15111511
__m128i sum_b_i32x4 = _mm_add_epi32(_mm256_castsi256_si128(state_b->sum_i32x8),
@@ -1514,7 +1514,7 @@ NK_INTERNAL void nk_dot_through_i32_finalize_haswell_(
15141514
_mm256_extracti128_si256(state_c->sum_i32x8, 1));
15151515
__m128i sum_d_i32x4 = _mm_add_epi32(_mm256_castsi256_si128(state_d->sum_i32x8),
15161516
_mm256_extracti128_si256(state_d->sum_i32x8, 1));
1517-
// Step 2: Transpose 4×4 matrix
1517+
// Transpose 4×4 matrix
15181518
__m128i transpose_ab_low_i32x4 = _mm_unpacklo_epi32(sum_a_i32x4, sum_b_i32x4);
15191519
__m128i transpose_cd_low_i32x4 = _mm_unpacklo_epi32(sum_c_i32x4, sum_d_i32x4);
15201520
__m128i transpose_ab_high_i32x4 = _mm_unpackhi_epi32(sum_a_i32x4, sum_b_i32x4);
@@ -1523,7 +1523,7 @@ NK_INTERNAL void nk_dot_through_i32_finalize_haswell_(
15231523
__m128i sum_lane1_i32x4 = _mm_unpackhi_epi64(transpose_ab_low_i32x4, transpose_cd_low_i32x4);
15241524
__m128i sum_lane2_i32x4 = _mm_unpacklo_epi64(transpose_ab_high_i32x4, transpose_cd_high_i32x4);
15251525
__m128i sum_lane3_i32x4 = _mm_unpackhi_epi64(transpose_ab_high_i32x4, transpose_cd_high_i32x4);
1526-
// Step 3: Vertical sum and store as i32
1526+
// Vertical sum and store as i32
15271527
__m128i sum_i32x4 = _mm_add_epi32(_mm_add_epi32(sum_lane0_i32x4, sum_lane1_i32x4),
15281528
_mm_add_epi32(sum_lane2_i32x4, sum_lane3_i32x4));
15291529
result->xmm = sum_i32x4;

0 commit comments

Comments
 (0)