Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 32 additions & 10 deletions compression/test_util-inl.h
Original file line number Diff line number Diff line change
Expand Up @@ -105,8 +105,6 @@ void ForeachActivationType3(D d) {
template <typename MatT>
MatStorageT<MatT> GenerateMat(const Extents2D& extents, MatPadding padding,
ThreadingContext& ctx) {
gcpp::CompressWorkingSet ws;
ws.tls.resize(ctx.pools.MaxWorkers());
MatStorageT<float> raw("raw", extents, ctx.allocator, MatPadding::kPacked);
MatStorageT<MatT> compressed("mat", extents, ctx.allocator, padding);
const float scale = SfpStream::kMax / extents.Area();
Expand All @@ -119,11 +117,24 @@ MatStorageT<MatT> 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<MatT>()) {
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;
}
Expand All @@ -134,8 +145,6 @@ template <typename MatT>
MatStorageT<MatT> GenerateTransposedMat(const Extents2D extents,
MatPadding padding,
ThreadingContext& ctx) {
gcpp::CompressWorkingSet ws;
ws.tls.resize(ctx.pools.MaxWorkers());
MatStorageT<float> raw("raw", extents, ctx.allocator, MatPadding::kPacked);
MatStorageT<MatT> compressed("trans", extents, ctx.allocator, padding);
const float scale = SfpStream::kMax / extents.Area();
Expand All @@ -148,11 +157,24 @@ MatStorageT<MatT> 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<MatT>()) {
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;
Expand Down
3 changes: 1 addition & 2 deletions compression/types.h
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
58 changes: 12 additions & 46 deletions gemma/weights.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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);
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -597,52 +609,6 @@ void T5GemmaDecoderLayerWeightsPtrs::Fixup(std::vector<MatOwner>& mat_owners,
self_qkv_einsum_w2);
}

static void HWY_MAYBE_UNUSED InitAttWeightsNUQ(
const LayerConfig& layer_config, MatPtrT<NuqStream>& attn_vec_einsum_w,
MatPtrT<NuqStream>& att_weights, std::vector<MatOwner>& 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<float[]> attn_vec_einsum_w_tmp =
hwy::AllocateAligned<float>(model_dim * heads * qkv_dim);
hwy::AlignedFreeUniquePtr<float[]> att_weights_tmp =
hwy::AllocateAligned<float>(model_dim * heads * qkv_dim);

const hwy::HWY_NAMESPACE::ScalableTag<float> 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) {
Expand Down
16 changes: 16 additions & 0 deletions ops/matmul_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -261,6 +261,22 @@ void TestAllMatMul() {
TestMatMul<F32, SFP>(256, 256, 256, /*add=*/false, env, __LINE__);
TestMatMul<BF16, SFP>(256, 256, 256, /*add=*/true, env, __LINE__);

#if GEMMA_ENABLE_NUQ
using NUQ = NuqStream;
TestMatMul<F32, NUQ>(256, 256, 256, /*add=*/false, env, __LINE__);
TestMatMul<BF16, NUQ>(256, 256, 256, /*add=*/true, env, __LINE__);
TestMatMul<F32, NUQ>(31, 128, 32, /*add=*/false, env, __LINE__);
TestMatMul<BF16, NUQ>(29, 128, 32, /*add=*/true, env, __LINE__);
TestMatMul<F32, NUQ>(4, 128, 32, /*add=*/true, env, __LINE__);
TestMatMul<BF16, NUQ>(4, 128, 32, /*add=*/false, env, __LINE__);
TestMatMul<F32, NUQ>(3, 128, 32, /*add=*/false, env, __LINE__);
TestMatMul<BF16, NUQ>(3, 128, 32, /*add=*/true, env, __LINE__);
TestMatMul<F32, NUQ>(2, 128, 64, /*add=*/true, env, __LINE__);
TestMatMul<BF16, NUQ>(2, 128, 64, /*add=*/false, env, __LINE__);
TestMatMul<F32, NUQ>(1, 128, 32, /*add=*/false, env, __LINE__);
TestMatMul<BF16, NUQ>(1, 128, 32, /*add=*/true, env, __LINE__);
#endif

// Non-vector-multiple K.
TestMatMul<F32, BF16>(128, 258, 128, /*add=*/true, env, __LINE__);
TestMatMul<BF16, BF16>(128, 258, 128, /*add=*/true, env, __LINE__);
Expand Down
Loading