diff --git a/BUILD.bazel b/BUILD.bazel index 47730efc..1a890bc1 100644 --- a/BUILD.bazel +++ b/BUILD.bazel @@ -156,7 +156,7 @@ cc_test( ":test_util", ":threading_context", ":weights", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_main", # buildcleaner: keep "//compression:compress", "//compression:types", "@highway//:hwy", @@ -185,10 +185,12 @@ cc_test( cc_library( name = "test_util", + testonly = True, hdrs = ["util/test_util.h"], deps = [ ":basics", ":mat", + "//testing/base/public:gunit_for_library_testonly", "@highway//:hwy", "@highway//:hwy_test_util", "@highway//:nanobenchmark", @@ -214,7 +216,7 @@ cc_test( srcs = ["gemma/configs_test.cc"], deps = [ ":configs", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_main", # buildcleaner: keep "//compression:types", "//io:fields", ], @@ -360,7 +362,7 @@ cc_test( ":mat", ":tensor_info", ":weights", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_main", # buildcleaner: keep "//compression:compress", "@highway//:hwy", # aligned_allocator.h ], @@ -526,7 +528,7 @@ cc_test( ":ops", ":test_util", ":threading_context", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_main", # buildcleaner: keep "//compression:compress", "//compression:test_util", "@highway//:hwy", @@ -557,7 +559,7 @@ cc_test( ":query", ":test_util", ":threading_context", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_main", # buildcleaner: keep "//compression:test_util", "//compression:types", "@highway//:hwy", @@ -583,7 +585,7 @@ cc_test( ":matmul_static", ":ops", ":threading_context", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_main", # buildcleaner: keep "//compression:compress", "//compression:test_util", "@highway//:hwy", @@ -881,7 +883,7 @@ cc_test( ":test_util", ":threading_context", ":weights", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_main", # buildcleaner: keep "//compression:compress", "//compression:types", "@highway//:hwy", @@ -965,8 +967,8 @@ cc_test( ":benchmark_helper", ":configs", ":gemma_lib", - "@googletest//:gtest_main", # buildcleaner: keep - "//io", + ":test_util", + "//testing/base/public:gunit_for_library_testonly", # buildcleaner: keep "@highway//:hwy", "@highway//:hwy_test_util", ], @@ -1006,7 +1008,8 @@ cc_test( deps = [ ":benchmark_helper", ":gemma_lib", - "@googletest//:gtest_main", # buildcleaner: keep + ":test_util", + "//testing/base/public:gunit_main", # buildcleaner: keep "@highway//:hwy", "@highway//:hwy_test_util", "@highway//:nanobenchmark", @@ -1039,7 +1042,8 @@ cc_test( ":benchmark_helper", ":configs", ":gemma_lib", - "@googletest//:gtest_main", # buildcleaner: keep + ":test_util", + "//testing/base/public:gunit_for_library_testonly", # buildcleaner: keep "//io", "@highway//:abort_header_only", "@highway//:hwy_test_util", @@ -1104,6 +1108,7 @@ cc_binary( deps = [ ":benchmark_helper", "@google_benchmark//:benchmark", + "//io", "@highway//:hwy", # base.h ], ) diff --git a/compression/BUILD.bazel b/compression/BUILD.bazel index feec5266..5e951f63 100644 --- a/compression/BUILD.bazel +++ b/compression/BUILD.bazel @@ -44,7 +44,7 @@ cc_test( srcs = ["distortion_test.cc"], deps = [ ":distortion", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_main", # buildcleaner: keep "//:test_util", "@highway//:hwy_test_util", "@highway//:nanobenchmark", # Unpredictable1 @@ -112,7 +112,7 @@ cc_test( tags = ["hwy_ops_test"], deps = [ ":int", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_main", # buildcleaner: keep "//:test_util", "@highway//:hwy", "@highway//:hwy_test_util", @@ -132,7 +132,7 @@ cc_test( deps = [ ":compress", ":q4_0", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_main", # buildcleaner: keep "//:test_util", "@highway//:hwy", "@highway//:hwy_test_util", @@ -141,6 +141,7 @@ cc_test( cc_library( name = "test_util", + testonly = True, textual_hdrs = [ "test_util-inl.h", ], @@ -165,7 +166,7 @@ cc_test( deps = [ ":compress", ":distortion", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_main", # buildcleaner: keep "//:test_util", "@highway//:hwy", "@highway//:hwy_test_util", @@ -185,7 +186,7 @@ cc_test( deps = [ ":distortion", ":nuq", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_main", # buildcleaner: keep "//:test_util", "@highway//:hwy", "@highway//:hwy_test_util", @@ -230,7 +231,7 @@ cc_test( ":compress", ":distortion", ":test_util", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_main", # buildcleaner: keep "//:test_util", "//:threading_context", "@highway//:hwy", diff --git a/evals/benchmark.cc b/evals/benchmark.cc index b142fe83..f03d3aff 100644 --- a/evals/benchmark.cc +++ b/evals/benchmark.cc @@ -130,6 +130,8 @@ int BenchmarkTriviaQA(GemmaEnv& env, const Path& json_file, } // namespace gcpp int main(int argc, char** argv) { + gcpp::InternalInit(); + gcpp::ConsumedArgs consumed(argc, argv); gcpp::GemmaArgs args(argc, argv, consumed); gcpp::BenchmarkArgs benchmark_args(argc, argv, consumed); diff --git a/evals/benchmark_helper.cc b/evals/benchmark_helper.cc index f7b74381..eabf351d 100644 --- a/evals/benchmark_helper.cc +++ b/evals/benchmark_helper.cc @@ -37,10 +37,7 @@ namespace gcpp { GemmaEnv::GemmaEnv(const GemmaArgs& args) - : initializer_value_(gcpp::InternalInit()), - ctx_(args.threading), - env_(ctx_), - gemma_(args, ctx_) { + : ctx_(args.threading), env_(ctx_), gemma_(args, ctx_) { const ModelConfig& config = gemma_.Config(); if (args.inference.verbosity >= 2) { diff --git a/evals/benchmark_helper.h b/evals/benchmark_helper.h index 85f0d21c..60447f40 100644 --- a/evals/benchmark_helper.h +++ b/evals/benchmark_helper.h @@ -125,8 +125,6 @@ class GemmaEnv { MatMulEnv& MutableEnv() { return env_; } private: - // This is used to ensure that InternalInit is called before anything else. - int initializer_value_ = 0; ThreadingContext ctx_; MatMulEnv env_; Gemma gemma_; diff --git a/evals/benchmarks.cc b/evals/benchmarks.cc index f44c62b6..f51bea8c 100644 --- a/evals/benchmarks.cc +++ b/evals/benchmarks.cc @@ -21,6 +21,7 @@ #include "benchmark/benchmark.h" #include "evals/benchmark_helper.h" #include "evals/prompts.h" +#include "io/io.h" namespace gcpp { @@ -98,6 +99,8 @@ BENCHMARK(BM_coding_prompt) ->UseRealTime(); int main(int argc, char** argv) { + gcpp::InternalInit(); + gcpp::ConsumedArgs consumed(argc, argv); gcpp::GemmaArgs args(argc, argv, consumed); consumed.AbortIfUnconsumed(); diff --git a/evals/debug_prompt.cc b/evals/debug_prompt.cc index a6cf8c48..ef0838a9 100644 --- a/evals/debug_prompt.cc +++ b/evals/debug_prompt.cc @@ -53,6 +53,8 @@ class PromptArgs : public ArgsBase { }; int Run(int argc, char** argv) { + InternalInit(); + ConsumedArgs consumed(argc, argv); const GemmaArgs args(argc, argv, consumed); const PromptArgs prompt_args(argc, argv, consumed); diff --git a/evals/gemma_batch_bench.cc b/evals/gemma_batch_bench.cc index ea5e9793..a2eecf14 100644 --- a/evals/gemma_batch_bench.cc +++ b/evals/gemma_batch_bench.cc @@ -21,10 +21,10 @@ #include "evals/benchmark_helper.h" #include "gemma/gemma.h" +#include "util/test_util.h" #include "hwy/base.h" #include "hwy/nanobenchmark.h" #include "hwy/profiler.h" -#include "hwy/tests/hwy_gtest.h" namespace gcpp { namespace { @@ -144,6 +144,8 @@ TEST_F(GemmaBatchBench, RandomQuestionsBatched) { } // namespace gcpp int main(int argc, char** argv) { + gcpp::InternalInitTest(); + fprintf(stderr, "GemmaEnv setup..\n"); gcpp::ConsumedArgs consumed(argc, argv); gcpp::GemmaArgs args(argc, argv, consumed); @@ -152,7 +154,5 @@ int main(int argc, char** argv) { gcpp::GemmaEnv env(args); gcpp::s_env = &env; - testing::InitGoogleTest(&argc, argv); - return RUN_ALL_TESTS(); } diff --git a/evals/gemma_test.cc b/evals/gemma_test.cc index a581561d..17ff78a0 100644 --- a/evals/gemma_test.cc +++ b/evals/gemma_test.cc @@ -22,16 +22,16 @@ #include "evals/benchmark_helper.h" #include "gemma/configs.h" +#include "util/test_util.h" #include "hwy/base.h" -#include "hwy/tests/hwy_gtest.h" // This test can be run manually with the downloaded gemma weights. // To run the test, pass the following flags: // --model --tokenizer --weights // or just use the single-file weights file with --weights . // It should pass for the following models: -// Gemma1: 2b-it (v1 and v1.1), 7b-it (v1 and v1.1), gr2b-it, -// Gemma2: gemma2-2b-it, 9b-it, 27b-it, +// Gemma2: gemma2-2b-it, 9b-it, 27b-it +// Gemma3: gemma3-270m-it namespace gcpp { namespace { @@ -183,7 +183,7 @@ TEST_F(GemmaTest, CrossEntropySmall) { } // namespace gcpp int main(int argc, char** argv) { - testing::InitGoogleTest(&argc, argv); + gcpp::InternalInitTest(); gcpp::GemmaTest::InitEnv(argc, argv); int ret = RUN_ALL_TESTS(); gcpp::GemmaTest::DeleteEnv(); diff --git a/evals/run_mmlu.cc b/evals/run_mmlu.cc index 8cbad072..66044397 100644 --- a/evals/run_mmlu.cc +++ b/evals/run_mmlu.cc @@ -152,6 +152,8 @@ void Run(GemmaEnv& env, JsonArgs& json) { } // namespace gcpp int main(int argc, char** argv) { + gcpp::InternalInit(); + { PROFILER_ZONE("Startup.all"); gcpp::ConsumedArgs consumed(argc, argv); diff --git a/evals/wheat_from_chaff_test.cc b/evals/wheat_from_chaff_test.cc index 8981beb4..8fcf4a9b 100644 --- a/evals/wheat_from_chaff_test.cc +++ b/evals/wheat_from_chaff_test.cc @@ -23,8 +23,8 @@ #include "gemma/configs.h" #include "gemma/gemma.h" #include "io/io.h" +#include "util/test_util.h" #include "hwy/base.h" -#include "hwy/tests/hwy_gtest.h" // This test can be run manually with the downloaded gemma weights. // To run the test, pass the following flags: @@ -187,7 +187,7 @@ TEST_F(GemmaTest, WheatFromChaff) { } // namespace gcpp int main(int argc, char** argv) { - testing::InitGoogleTest(&argc, argv); + gcpp::InternalInitTest(); gcpp::GemmaTest::InitEnv(argc, argv); int ret = RUN_ALL_TESTS(); gcpp::GemmaTest::DeleteEnv(); diff --git a/io/BUILD.bazel b/io/BUILD.bazel index 3d88ef90..0dc53445 100644 --- a/io/BUILD.bazel +++ b/io/BUILD.bazel @@ -89,7 +89,7 @@ cc_test( deps = [ ":blob_store", ":io", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_main", # buildcleaner: keep "//:basics", "//:threading_context", "@highway//:hwy_test_util", diff --git a/io/io.cc b/io/io.cc index 2f479b21..e9f27c2f 100644 --- a/io/io.cc +++ b/io/io.cc @@ -237,8 +237,7 @@ bool IOBatch::Add(void* mem, size_t bytes) { } int InternalInit() { - // currently unused, except for init list ordering in GemmaEnv. - return 0; + return 0; // currently unused } uint64_t IOBatch::Read(const File& file) const { diff --git a/io/migrate_weights.cc b/io/migrate_weights.cc index beb268e2..e68f94ed 100644 --- a/io/migrate_weights.cc +++ b/io/migrate_weights.cc @@ -40,6 +40,8 @@ struct WriterArgs : public ArgsBase { } // namespace gcpp int main(int argc, char** argv) { + gcpp::InternalInit(); + gcpp::ConsumedArgs consumed(argc, argv); gcpp::GemmaArgs args(argc, argv, consumed); gcpp::WriterArgs writer_args(argc, argv, consumed); diff --git a/paligemma/BUILD.bazel b/paligemma/BUILD.bazel index b749e05d..e04d8439 100644 --- a/paligemma/BUILD.bazel +++ b/paligemma/BUILD.bazel @@ -60,11 +60,12 @@ cc_test( ], deps = [ ":paligemma_helper", - "@googletest//:gtest_main", # buildcleaner: keep + "//testing/base/public:gunit_for_library_testonly", # buildcleaner: keep "//:allocator", "//:benchmark_helper", "//:configs", "//:gemma_lib", + "//:test_util", "@highway//:hwy_test_util", ], ) diff --git a/paligemma/paligemma_test.cc b/paligemma/paligemma_test.cc index 7bfd78cd..a438dd3a 100644 --- a/paligemma/paligemma_test.cc +++ b/paligemma/paligemma_test.cc @@ -23,7 +23,7 @@ #include "gemma/gemma.h" #include "paligemma/paligemma_helper.h" #include "util/allocator.h" -#include "hwy/tests/hwy_gtest.h" +#include "util/test_util.h" // This test can be run manually with the downloaded PaliGemma weights. // It should pass for `paligemma-3b-mix-224` and `paligemma2-3b-pt-448`. @@ -70,7 +70,7 @@ TEST_F(PaliGemmaTest, QueryObjects) { } // namespace gcpp int main(int argc, char** argv) { - testing::InitGoogleTest(&argc, argv); + gcpp::InternalInitTest(); gcpp::ConsumedArgs consumed(argc, argv); gcpp::GemmaArgs args(argc, argv, consumed); diff --git a/util/test_util.h b/util/test_util.h index 443990f8..c1a812c3 100644 --- a/util/test_util.h +++ b/util/test_util.h @@ -23,6 +23,7 @@ #include #include +#include "gtest/gtest.h" #include "util/basics.h" // RngStream #include "util/mat.h" #include "hwy/base.h" @@ -35,6 +36,10 @@ namespace gcpp { +inline int InternalInitTest() { + return 0; // currently unused +} + // Excludes outliers; we might not have enough samples for a reliable mode. HWY_INLINE double TrimmedMean(double* seconds, size_t num) { std::sort(seconds, seconds + num);