Skip to content

Commit b57557a

Browse files
pcullitoncopybara-github
authored andcommitted
Internal change.
PiperOrigin-RevId: 979960879
1 parent 8073ef9 commit b57557a

13 files changed

Lines changed: 580 additions & 40 deletions

‎BUILD.bazel‎

Lines changed: 3 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1056,7 +1056,6 @@ cc_test(
10561056
cc_test(
10571057
name = "model_health_test",
10581058
srcs = ["evals/model_health_test.cc"],
1059-
linkstatic = True,
10601059
# Requires model files
10611060
tags = [
10621061
"local",
@@ -1068,7 +1067,8 @@ cc_test(
10681067
":configs",
10691068
":cross_entropy",
10701069
":gemma_lib",
1071-
"@googletest//:gtest_main", # buildcleaner: keep
1070+
":test_util",
1071+
"//testing/base/public:gunit_for_library_testonly", # buildcleaner: keep
10721072
"@highway//:hwy",
10731073
"@highway//:hwy_test_util",
10741074
],
@@ -1077,7 +1077,6 @@ cc_test(
10771077
cc_test(
10781078
name = "gemma_batch_bench",
10791079
srcs = ["evals/gemma_batch_bench.cc"],
1080-
linkstatic = True,
10811080
# Requires model files
10821081
tags = [
10831082
"local",
@@ -1088,7 +1087,7 @@ cc_test(
10881087
":benchmark_helper",
10891088
":gemma_lib",
10901089
":test_util",
1091-
"@googletest//:gtest_main", # buildcleaner: keep
1090+
"//testing/base/public:gunit_for_library_testonly", # buildcleaner: keep
10921091
"@highway//:hwy",
10931092
"@highway//:hwy_test_util",
10941093
"@highway//:nanobenchmark",

‎compression/compress.cc‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515

1616
#include "compression/compress.h"
1717

18+
#include <cmath>
1819
#include <stddef.h>
1920
#include <stdint.h>
2021

‎compression/python/compression_clif_aux.cc‎

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -85,8 +85,10 @@ class SbsWriterImpl : public ISbsWriter {
8585
}
8686

8787
HWY_ASSERT(weights.size() == mat.Extents().Area());
88-
Compress(weights.data(), weights.size(), working_set_, mat.Span(),
89-
/*packed_ofs=*/0, ctx_);
88+
{
89+
Compress(weights.data(), weights.size(), working_set_, mat.Span(),
90+
/*packed_ofs=*/0, ctx_);
91+
}
9092
writer_.Add(name, mat.Packed(), mat.PackedBytes());
9193
}
9294

‎compression/types.h‎

Lines changed: 13 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -269,7 +269,8 @@ constexpr bool IsMxFp4Stream() {
269269
template <typename Packed>
270270
constexpr bool IsPacked() {
271271
return IsNuqStream<Packed>() || IsI8Stream<Packed>() ||
272-
IsQ4_0Stream<Packed>() || IsMxFp4Stream<Packed>();
272+
IsQ4_0Stream<Packed>() || IsMxFp4Stream<Packed>()
273+
;
273274
}
274275

275276
template <typename Packed>
@@ -294,12 +295,14 @@ enum class Type {
294295
kInt8,
295296
kQ4_0,
296297
kMXFP4,
298+
kReserved14,
297299
};
298300
// These are used in `ModelConfig.Specifier`, hence the strings will not
299301
// change, though new ones may be added.
300302
static constexpr const char* kTypeStrings[] = {
301-
"unknown", "f32", "bf16", "sfp", "nuq", "f64", "u32",
302-
"u64", "i8", "u16", "u8", "int8", "q4_0", "mxfp4"};
303+
"unknown", "f32", "bf16", "sfp", "nuq", "f64", "u32", "u64",
304+
"i8", "u16", "u8", "int8", "q4_0", "mxfp4", "reserved14"
305+
};
303306
static constexpr size_t kNumTypes =
304307
sizeof(kTypeStrings) / sizeof(kTypeStrings[0]);
305308
static constexpr size_t kTypeBits[] = {
@@ -317,6 +320,7 @@ static constexpr size_t kTypeBits[] = {
317320
8 * sizeof(int8_t),
318321
4 /* Q4_0Stream, actually 4.5 */,
319322
4 /* MxFp4Stream, actually 4.25 */,
323+
0 /* reserved */,
320324
};
321325

322326
static inline bool EnumValid(Type type) {
@@ -376,17 +380,20 @@ constexpr bool IsCompressed() {
376380
hwy::IsSame<hwy::RemoveCvRef<Packed>, NuqStream>() ||
377381
hwy::IsSame<hwy::RemoveCvRef<Packed>, I8Stream>() ||
378382
hwy::IsSame<hwy::RemoveCvRef<Packed>, Q4_0Stream>() ||
379-
hwy::IsSame<hwy::RemoveCvRef<Packed>, MxFp4Stream>();
383+
hwy::IsSame<hwy::RemoveCvRef<Packed>, MxFp4Stream>()
384+
;
380385
}
381386

382387
static inline bool IsCompressed(Type type) {
383388
return type == Type::kSFP || type == Type::kNUQ || type == Type::kI8 ||
384-
type == Type::kQ4_0 || type == Type::kMXFP4;
389+
type == Type::kQ4_0 || type == Type::kMXFP4
390+
;
385391
}
386392

387393
static inline bool IsPacked(Type type) {
388394
return type == Type::kNUQ || type == Type::kI8 || type == Type::kQ4_0 ||
389-
type == Type::kMXFP4;
395+
type == Type::kMXFP4
396+
;
390397
}
391398

392399
static inline bool SupportsPointerArithmetic(Type type) {

‎evals/model_health_test.cc‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -44,6 +44,7 @@
4444
#include "evals/cross_entropy.h"
4545
#include "gemma/configs.h"
4646
#include "gemma/gemma.h"
47+
#include "util/test_util.h"
4748
#include "hwy/base.h"
4849
#include "hwy/tests/hwy_gtest.h"
4950

@@ -485,7 +486,7 @@ TEST_F(ModelHealthTest, DeterministicGeneration) {
485486
} // namespace gcpp
486487

487488
int main(int argc, char** argv) {
488-
testing::InitGoogleTest(&argc, argv);
489+
gcpp::InternalInitTest();
489490
gcpp::ModelHealthTest::InitEnv(argc, argv);
490491
int ret = RUN_ALL_TESTS();
491492
gcpp::ModelHealthTest::DeleteEnv();

‎gemma/activations.h‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -827,7 +827,8 @@ struct Activations {
827827
config.model_dim, allocator)),
828828
mla_o_in(MatFactory("mla_o_in",
829829
mla_dims.o_in_dim > 0 ? batch_size : 0,
830-
mla_dims.o_in_dim, allocator)) {
830+
mla_dims.o_in_dim, allocator))
831+
{
831832
moe_C1.AllocateAndAttachRowPtrs(row_ptrs);
832833
moe_C2.AllocateAndAttachRowPtrs(row_ptrs);
833834
ffw_expert_in.AllocateAndAttachRowPtrs(row_ptrs);

‎gemma/configs.cc‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -481,7 +481,7 @@ static LayerConfig LayerConfigGemma4_26B_MoE_LM(size_t model_dim) {
481481
static ModelConfig ConfigGemma4_26B_MoE() {
482482
ModelConfig config = ConfigBaseGemmaV4();
483483
config.display_name = "Gemma4_26B_MoE";
484-
config.final_cap = 0.0f;
484+
config.final_cap = 30.0f;
485485
config.att_cap = 0.0f;
486486
config.model = Model::GEMMA4_26B_MOE;
487487
config.wrapping = PromptWrapping::GEMMA_IT;

‎gemma/gemma4_moe.cc‎

Lines changed: 26 additions & 19 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,7 @@
4848
#include "gemma/attention.h" // includes highway.h
4949
#include "gemma/tiled_attention.h"
5050
#include "gemma/gemma-inl.h"
51+
#include "ops/fast_ops-inl.h"
5152
#include "ops/ops-inl.h"
5253

5354
HWY_BEFORE_NAMESPACE();
@@ -239,6 +240,8 @@ struct Gemma4MoE {
239240
activations.s_expert_in.Notify(layer.layer_idx, tmp_in, env.ctx, 0,
240241
cluster_idx, parallelism);
241242

243+
const size_t expert_ff_hidden_dim =
244+
layer.moe_gating_einsum_w1[expert_idx].Rows();
242245
const ActivationType activation = layer.layer_config.activation;
243246
MatPtrT<BF16>& C1 = per_cluster.moe_C1;
244247
C1.OverrideRows(expert_size);
@@ -275,19 +278,20 @@ struct Gemma4MoE {
275278
cluster_idx, parallelism);
276279

277280
// Hidden layer -> output layer.
278-
size_t expert_ff_hidden_dim = layer.moe_gating_einsum_w1[expert_idx].Rows();
279-
MatStorageT<BF16> C1_narrow("C1_n",
280-
Extents2D(expert_size, expert_ff_hidden_dim),
281-
env.ctx.allocator, MatPadding::kOdd);
282-
for (size_t i = 0; i < expert_size; ++i) {
283-
memcpy(C1_narrow.Row(i), C1.Row(i), expert_ff_hidden_dim * sizeof(BF16));
281+
{
282+
MatStorageT<BF16> C1_narrow("C1_n",
283+
Extents2D(expert_size, expert_ff_hidden_dim),
284+
env.ctx.allocator, MatPadding::kOdd);
285+
for (size_t i = 0; i < expert_size; ++i) {
286+
memcpy(C1_narrow.Row(i), C1.Row(i),
287+
expert_ff_hidden_dim * sizeof(BF16));
288+
}
289+
CallMatMul(C1_narrow, layer.moe_linear_w[expert_idx],
290+
/*add=*/nullptr, env, expert_out, options);
284291
}
285292

286-
287-
CallMatMul(C1_narrow, layer.moe_linear_w[expert_idx],
288-
/*add=*/nullptr, env, expert_out, options);
289-
290-
if (layer.p_expert_sc.HasPtr()) {
293+
if (layer.p_expert_sc.HasPtr()
294+
) {
291295
const MatPtrT<BF16> sc_mat(layer.p_expert_sc);
292296
const BF16* sc_data = sc_mat.Row(0);
293297
const float expert_scale =
@@ -466,9 +470,12 @@ void Gemma4MoETransformerLayer(size_t num_tokens, size_t layer_idx,
466470
/*is_attention=*/true, env.ctx);
467471

468472
// Dual-Path FFW
469-
pre_norm(
470-
layer.pre_ffw2_ns.HasPtr() ? layer.pre_ffw2_ns : layer.pre_ffw_norm_scale,
471-
activations.pre_ffw_rms_out);
473+
const MatPtr& shared_norm = layer.pre_ffw2_ns.HasPtr() ? layer.pre_ffw2_ns : layer.pre_ffw_norm_scale;
474+
const MatPtr& moe_norm = layer.pre_ffw_norm_scale;
475+
const MatPtr& shared_post_norm = layer.post_ffw2_ns;
476+
const MatPtr& moe_post_norm = layer.post_ffw1_ns;
477+
478+
pre_norm(shared_norm, activations.pre_ffw_rms_out);
472479

473480
// Shared MLP Path
474481
FFWNoVit(layer, activations, env); // writes to activations.ffw_out
@@ -482,17 +489,17 @@ void Gemma4MoETransformerLayer(size_t num_tokens, size_t layer_idx,
482489
}
483490
}
484491

485-
if (layer.post_ffw2_ns.HasPtr()) {
486-
rms_norm_inplace(layer.post_ffw2_ns, activations.attention.att_sums);
492+
if (shared_post_norm.HasPtr()) {
493+
rms_norm_inplace(shared_post_norm, activations.attention.att_sums);
487494
}
488495

489496
// MoE Path
490-
pre_norm(layer.pre_ffw_norm_scale, activations.pre_ffw_rms_out);
497+
pre_norm(moe_norm, activations.pre_ffw_rms_out);
491498

492499
Gemma4MoE::MoEFFW(layer, activations, env); // writes to activations.ffw_out
493500

494-
if (layer.post_ffw1_ns.HasPtr()) {
495-
rms_norm_inplace(layer.post_ffw1_ns, activations.ffw_out);
501+
if (moe_post_norm.HasPtr()) {
502+
rms_norm_inplace(moe_post_norm, activations.ffw_out);
496503
}
497504

498505
// Combine & Final Norm (Fix for dual-path combination)

‎gemma/tensor_info.h‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -68,6 +68,9 @@ struct TensorInfo {
6868
// If false, then [10, 20, 30] -> [10*20, 30] and [30] -> [1, 30].
6969
// If true, then [10, 20, 30] -> [10, 20*30] and [30] -> [1, 30].
7070
bool cols_take_extra_dims = false;
71+
// Optional pre-computed scale (e.g. for kW2_UL weights from QAFT).
72+
// If > 0.0, used directly instead of re-estimating via ScaleWeightsW2UL.
73+
float scale = 0.0f;
7174
};
7275

7376
// Collapses/expands the tensor dims into 2D extents, which may be 0, 0 for

‎gemma/weights.cc‎

Lines changed: 17 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@
2020
#include <stdio.h>
2121
#include <stdlib.h>
2222

23+
#include <cstring>
2324
#include <mutex> // NOLINT
2425
#include <string>
2526
#include <vector>
@@ -129,7 +130,10 @@ static void SplitPackedMatrix(MatPtr& parent, size_t split_row, MatPtr& w1,
129130
void LayerWeightsPtrs::SplitW1() {
130131
// Used for Gemma layers; FFWVit uses different tensors.
131132
if (layer_config.type == LayerAttentionType::kVit) return;
132-
if (layer_config.IsMoE()) return;
133+
if (layer_config.IsMoE() && !gating_einsum_w.HasPtr() &&
134+
!gating_einsum_w1.HasPtr()) {
135+
return;
136+
}
133137

134138
// Files have both or neither of w1 and w2.
135139
HWY_ASSERT(gating_einsum_w1.HasPtr() == gating_einsum_w2.HasPtr());
@@ -543,6 +547,18 @@ void LayerWeightsPtrs::Fixup(Model model, std::vector<MatOwner>& mat_owners,
543547
const size_t elem_bytes = qkv_einsum_w2.ElementBytes();
544548
const size_t old_row_bytes = old_stride * elem_bytes;
545549
const size_t kv_heads = layer_config.kv_heads;
550+
const size_t qkv_dim = layer_config.qkv_dim;
551+
552+
// In Gemma 4 global layers, attention_k_eq_v is true (K0 == V0).
553+
// If already interleaved by exporter: [K0, V0, K1, V1], slice 0 == slice 1.
554+
// If not interleaved: [K0, K1, V0, V1], slice 0 (K0) != slice 1 (K1).
555+
const uint8_t* slice0 = qkv_einsum_w2.RowBytes(0);
556+
const uint8_t* slice1 = qkv_einsum_w2.RowBytes(qkv_dim);
557+
if (std::memcmp(slice0, slice1, old_row_bytes) == 0) {
558+
// Exporter already emitted interleaved layout; skip fixup.
559+
return;
560+
}
561+
546562
const size_t total_bytes = qkv_einsum_w2.Rows() * old_row_bytes;
547563
hwy::AlignedFreeUniquePtr<uint8_t[]> tmp =
548564
hwy::AllocateAligned<uint8_t>(total_bytes);
@@ -556,7 +572,6 @@ void LayerWeightsPtrs::Fixup(Model model, std::vector<MatOwner>& mat_owners,
556572
}
557573

558574
const size_t new_row_bytes = qkv_einsum_w2.Cols() * elem_bytes;
559-
const size_t qkv_dim = layer_config.qkv_dim;
560575
const uint8_t* src_ptr = tmp.get();
561576
for (size_t i = 0; i < kv_heads; ++i) {
562577
for (size_t row = 0; row < qkv_dim; ++row) {

0 commit comments

Comments
 (0)