From 550bf2d508d08dc3ba116a5ae34e20955c1c5d65 Mon Sep 17 00:00:00 2001 From: Phil Culliton Date: Fri, 31 Jul 2026 17:28:00 -0700 Subject: [PATCH] Internal change. PiperOrigin-RevId: 957408238 --- compression/test_util-inl.h | 42 ++++++++++++++++++++------- compression/types.h | 3 +- gemma/weights.cc | 58 ++++++++----------------------------- ops/matmul_test.cc | 16 ++++++++++ 4 files changed, 61 insertions(+), 58 deletions(-) diff --git a/compression/test_util-inl.h b/compression/test_util-inl.h index bb2fadb0..187e1456 100644 --- a/compression/test_util-inl.h +++ b/compression/test_util-inl.h @@ -105,8 +105,6 @@ void ForeachActivationType3(D d) { template MatStorageT GenerateMat(const Extents2D& extents, MatPadding padding, ThreadingContext& ctx) { - gcpp::CompressWorkingSet ws; - ws.tls.resize(ctx.pools.MaxWorkers()); MatStorageT raw("raw", extents, ctx.allocator, MatPadding::kPacked); MatStorageT compressed("mat", extents, ctx.allocator, padding); const float scale = SfpStream::kMax / extents.Area(); @@ -119,11 +117,24 @@ MatStorageT GenerateMat(const Extents2D& extents, MatPadding padding, f = -f; // Also generate some negative values. row[c] = f; } - Compress(raw.Row(r), raw.Cols(), ws.tls[thread], - MakeSpan(compressed.Row(r), extents.cols), - /*packed_ofs=*/0); }); + if constexpr (IsCompressed()) { + gcpp::CompressWorkingSet ws; + ws.tls.resize(ctx.pools.MaxWorkers()); + Compress(raw.Row(0), extents.Area(), ws, compressed.Span(), + /*packed_ofs=*/0, ctx); + } else { + gcpp::CompressWorkingSet ws; + ws.tls.resize(ctx.pools.MaxWorkers()); + ParallelFor(Parallelism::kFlat, extents.rows, ctx, /*cluster_idx=*/0, + Callers::kTest, [&](size_t r, size_t thread) { + Compress(raw.Row(r), raw.Cols(), ws.tls[thread], + MakeSpan(compressed.Row(r), extents.cols), + /*packed_ofs=*/0); + }); + } + compressed.SetScale(0.6f); // Arbitrary value, different from 1. return compressed; } @@ -134,8 +145,6 @@ template MatStorageT GenerateTransposedMat(const Extents2D extents, MatPadding padding, ThreadingContext& ctx) { - gcpp::CompressWorkingSet ws; - ws.tls.resize(ctx.pools.MaxWorkers()); MatStorageT raw("raw", extents, ctx.allocator, MatPadding::kPacked); MatStorageT compressed("trans", extents, ctx.allocator, padding); const float scale = SfpStream::kMax / extents.Area(); @@ -148,11 +157,24 @@ MatStorageT GenerateTransposedMat(const Extents2D extents, f = -f; // Also generate some negative values. row[c] = f; } - Compress(raw.Row(r), raw.Cols(), ws.tls[thread], - MakeSpan(compressed.Row(r), extents.cols), - /*packed_ofs=*/0); }); + if constexpr (IsCompressed()) { + gcpp::CompressWorkingSet ws; + ws.tls.resize(ctx.pools.MaxWorkers()); + Compress(raw.Row(0), extents.Area(), ws, compressed.Span(), + /*packed_ofs=*/0, ctx); + } else { + gcpp::CompressWorkingSet ws; + ws.tls.resize(ctx.pools.MaxWorkers()); + ParallelFor(Parallelism::kFlat, extents.rows, ctx, /*cluster_idx=*/0, + Callers::kTest, [&](size_t r, size_t thread) { + Compress(raw.Row(r), raw.Cols(), ws.tls[thread], + MakeSpan(compressed.Row(r), extents.cols), + /*packed_ofs=*/0); + }); + } + // Arbitrary value, different from 1, must match `GenerateMat`. compressed.SetScale(0.6f); return compressed; diff --git a/compression/types.h b/compression/types.h index 7a4f86cf..9454412f 100644 --- a/compression/types.h +++ b/compression/types.h @@ -59,9 +59,8 @@ namespace gcpp { #endif // GEMMA_DISABLED_TARGETS -// Only used in experiments, hence disable in default builds. #ifndef GEMMA_ENABLE_NUQ -#define GEMMA_ENABLE_NUQ 0 +#define GEMMA_ENABLE_NUQ 1 #endif // Switching Floating Point: a hybrid 8-bit float representation of bf16/f32 diff --git a/gemma/weights.cc b/gemma/weights.cc index 0ad8f473..9292a640 100644 --- a/gemma/weights.cc +++ b/gemma/weights.cc @@ -137,6 +137,12 @@ void LayerWeightsPtrs::SplitW1() { uint8_t* base_ptr = gating_einsum_w.RowBytes(0); gating_einsum_w1.SetPtr(base_ptr, stride); gating_einsum_w2.SetPtr(base_ptr + Q4_0Stream::PackedEnd(ff_hidden_dim * stride), stride); + } else if (gating_einsum_w.GetType() == Type::kNUQ) { + const size_t stride = gating_einsum_w.Stride(); + uint8_t* base_ptr = gating_einsum_w.RowBytes(0); + gating_einsum_w1.SetPtr(base_ptr, stride); + gating_einsum_w2.SetPtr( + base_ptr + NuqStream::PackedEnd(ff_hidden_dim * stride), stride); } else { const size_t stride = gating_einsum_w.Stride(); gating_einsum_w1.SetPtr(gating_einsum_w.RowBytes(0), stride); @@ -191,6 +197,12 @@ void LayerWeightsPtrs::SplitAttW1() { uint8_t* base_ptr = qkv_einsum_w.RowBytes(0); qkv_einsum_w1.SetPtr(base_ptr, stride); qkv_einsum_w2.SetPtr(base_ptr + Q4_0Stream::PackedEnd(w1_rows * stride), stride); + } else if (qkv_einsum_w.GetType() == Type::kNUQ) { + const size_t stride = qkv_einsum_w.Stride(); + uint8_t* base_ptr = qkv_einsum_w.RowBytes(0); + qkv_einsum_w1.SetPtr(base_ptr, stride); + qkv_einsum_w2.SetPtr(base_ptr + NuqStream::PackedEnd(w1_rows * stride), + stride); } else { const size_t stride = qkv_einsum_w.Stride(); qkv_einsum_w1.SetPtr(qkv_einsum_w.RowBytes(0), stride); @@ -597,52 +609,6 @@ void T5GemmaDecoderLayerWeightsPtrs::Fixup(std::vector& mat_owners, self_qkv_einsum_w2); } -static void HWY_MAYBE_UNUSED InitAttWeightsNUQ( - const LayerConfig& layer_config, MatPtrT& attn_vec_einsum_w, - MatPtrT& att_weights, std::vector& mat_owners, - ThreadingContext& ctx) { - if (!attn_vec_einsum_w.HasPtr()) return; - HWY_ASSERT(attn_vec_einsum_w.GetType() == Type::kNUQ); - - HWY_ASSERT(att_weights.HasPtr()); - att_weights.SetType(Type::kNUQ); - - const size_t model_dim = layer_config.model_dim; - const size_t heads = layer_config.heads; - const size_t qkv_dim = layer_config.qkv_dim; - - // Reshape [kHeads, kModelDim, kQKVDim] to [kModelDim, kHeads * kQKVDim]. - hwy::AlignedFreeUniquePtr attn_vec_einsum_w_tmp = - hwy::AllocateAligned(model_dim * heads * qkv_dim); - hwy::AlignedFreeUniquePtr att_weights_tmp = - hwy::AllocateAligned(model_dim * heads * qkv_dim); - - const hwy::HWY_NAMESPACE::ScalableTag df; - HWY_NAMESPACE::DecompressAndZeroPad(df, attn_vec_einsum_w.Span(), 0, - attn_vec_einsum_w_tmp.get(), - model_dim * heads * qkv_dim); - - for (size_t m = 0; m < model_dim; ++m) { - float* HWY_RESTRICT out_row = att_weights_tmp.get() + m * heads * qkv_dim; - for (size_t h = 0; h < heads; ++h) { - hwy::CopyBytes( - attn_vec_einsum_w_tmp.get() + h * model_dim * qkv_dim + m * qkv_dim, - out_row + h * qkv_dim, qkv_dim * sizeof(float)); - } - } - - CompressWorkingSet work; - HWY_NAMESPACE::Compress(att_weights_tmp.get(), model_dim * heads * qkv_dim, - work, att_weights.Span(), - /*packed_ofs=*/0, ctx); - - att_weights.SetScale(attn_vec_einsum_w.Scale()); -} - -static void HWY_MAYBE_UNUSED SplitW1NUQ(const LayerConfig& layer_config) { - // TODO(janwas): implement. -} - // Zero-initializes only the allocated tensors in `*this`. void WeightsPtrs::ZeroInit() { ForEachTensor(nullptr, nullptr, [](const TensorArgs& t) { diff --git a/ops/matmul_test.cc b/ops/matmul_test.cc index 5ae9e366..790948b9 100644 --- a/ops/matmul_test.cc +++ b/ops/matmul_test.cc @@ -261,6 +261,22 @@ void TestAllMatMul() { TestMatMul(256, 256, 256, /*add=*/false, env, __LINE__); TestMatMul(256, 256, 256, /*add=*/true, env, __LINE__); +#if GEMMA_ENABLE_NUQ + using NUQ = NuqStream; + TestMatMul(256, 256, 256, /*add=*/false, env, __LINE__); + TestMatMul(256, 256, 256, /*add=*/true, env, __LINE__); + TestMatMul(31, 128, 32, /*add=*/false, env, __LINE__); + TestMatMul(29, 128, 32, /*add=*/true, env, __LINE__); + TestMatMul(4, 128, 32, /*add=*/true, env, __LINE__); + TestMatMul(4, 128, 32, /*add=*/false, env, __LINE__); + TestMatMul(3, 128, 32, /*add=*/false, env, __LINE__); + TestMatMul(3, 128, 32, /*add=*/true, env, __LINE__); + TestMatMul(2, 128, 64, /*add=*/true, env, __LINE__); + TestMatMul(2, 128, 64, /*add=*/false, env, __LINE__); + TestMatMul(1, 128, 32, /*add=*/false, env, __LINE__); + TestMatMul(1, 128, 32, /*add=*/true, env, __LINE__); +#endif + // Non-vector-multiple K. TestMatMul(128, 258, 128, /*add=*/true, env, __LINE__); TestMatMul(128, 258, 128, /*add=*/true, env, __LINE__);