--- a/mlx/backend/cpu/quantized.cpp +++ b/mlx/backend/cpu/quantized.cpp @@ -1,4 +1,6 @@ // Copyright © 2023-2026 Apple Inc. + +#include #include "mlx/backend/common/quantized.h" #include "mlx/backend/common/unary.h" @@ -458,14 +460,15 @@ constexpr int pack_factor = get_pack_factor(bits, 8); constexpr int packs_in_group = group_size / pack_factor; + std::vector accum(N); for (int m = 0; m < M; m++) { const uint8_t* w_local = (const uint8_t*)w; const uint8_t* scales_local = scales; - std::fill(result, result + N, 0); + std::fill(accum.begin(), accum.end(), 0.0f); for (int k = 0; k < K; k++) { - T* result_local = result; + float* result_local = accum.data(); T xi = *x++; for (int n = 0; n < N; n += group_size) { @@ -483,6 +486,9 @@ w_local++; } } + } + for (int n = 0; n < N; n++) { + result[n] = static_cast(accum[n]); } result += N; }