Skip to content

Commit cab77f8

Browse files
jan-wassenbergcopybara-github
authored andcommitted
Improved timing for image tokens
Move to TimingInfo, extra newline before profiler PiperOrigin-RevId: 881943820
1 parent 70cb9cf commit cab77f8

7 files changed

Lines changed: 71 additions & 46 deletions

File tree

gemma/bindings/context.cc

Lines changed: 1 addition & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -29,7 +29,6 @@
2929
#include "util/threading.h"
3030
#include "util/threading_context.h"
3131
#include "hwy/profiler.h"
32-
#include "hwy/timer.h"
3332

3433
#ifdef _WIN32
3534
#include <Windows.h>
@@ -195,17 +194,10 @@ int GemmaContext::GenerateInternal(const char* prompt_string,
195194

196195
// Use the existing runtime_config defined earlier in the function.
197196
// RuntimeConfig runtime_config = { ... }; // This was already defined
198-
double image_tokens_start = hwy::platform::Now();
199197
// Pass the populated image object to GenerateImageTokens
200198
model.GenerateImageTokens(runtime_config,
201199
active_conversation->kv_cache->SeqLen(), image,
202-
image_tokens, matmul_env);
203-
double image_tokens_duration = hwy::platform::Now() - image_tokens_start;
204-
205-
ss.str("");
206-
ss << "\n\n[ Timing info ] Image token generation took: ";
207-
ss << static_cast<int>(image_tokens_duration * 1000) << " ms\n",
208-
LogDebug(ss.str().c_str());
200+
image_tokens, matmul_env, timing_info);
209201

210202
prompt = WrapAndTokenize(
211203
model.Tokenizer(), model.ChatTemplate(), model_config.wrapping,

gemma/gemma.cc

Lines changed: 27 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -605,6 +605,9 @@ static void GenerateT(const ModelConfig& config,
605605
config, runtime_config, weights, activations, qbatch, env, timing_info);
606606
// No-op if the profiler is disabled, but useful to separate prefill and
607607
// generate phases for profiling.
608+
if constexpr (PROFILER_ENABLED) {
609+
fprintf(stderr, "\n");
610+
}
608611
env.ctx.profiler.PrintResults();
609612

610613
hwy::BitSet4096<> non_eos; // indexed by qi
@@ -725,25 +728,33 @@ void GenerateBatchT(const ModelConfig& config,
725728
void GenerateImageTokensT(const ModelConfig& config,
726729
const RuntimeConfig& runtime_config, size_t seq_len,
727730
const WeightsPtrs& weights, const Image& image,
728-
ImageTokens& image_tokens, MatMulEnv& env) {
729-
GCPP_ZONE(env.ctx, hwy::Profiler::GlobalIdx(), Zones::kGenImageTokens);
730-
if (config.vit_config.layer_configs.empty()) {
731-
HWY_ABORT("Model does not support generating image tokens.");
732-
}
733-
RuntimeConfig prefill_runtime_config = runtime_config;
731+
ImageTokens& image_tokens, MatMulEnv& env,
732+
TimingInfo& timing_info) {
734733
const ModelConfig vit_config = GetVitConfig(config);
735734
const size_t num_tokens = vit_config.max_seq_len;
736-
prefill_runtime_config.prefill_tbatch_size =
737-
num_tokens / (vit_config.pool_dim * vit_config.pool_dim);
738-
Activations prefill_activations(runtime_config, vit_config, num_tokens,
739-
num_tokens, env.ctx, env.row_ptrs);
740-
// Weights are for the full PaliGemma model, not just the ViT part.
741-
PrefillVit(config, weights, prefill_runtime_config, image, image_tokens,
742-
prefill_activations, env);
735+
736+
timing_info.NotifyImageTokenStart();
737+
738+
{
739+
GCPP_ZONE(env.ctx, hwy::Profiler::GlobalIdx(), Zones::kGenImageTokens);
740+
if (config.vit_config.layer_configs.empty()) {
741+
HWY_ABORT("Model does not support generating image tokens.");
742+
}
743+
RuntimeConfig prefill_runtime_config = runtime_config;
744+
prefill_runtime_config.prefill_tbatch_size =
745+
num_tokens / (vit_config.pool_dim * vit_config.pool_dim);
746+
Activations prefill_activations(runtime_config, vit_config, num_tokens,
747+
num_tokens, env.ctx, env.row_ptrs);
748+
// Weights are for the full PaliGemma model, not just the ViT part.
749+
PrefillVit(config, weights, prefill_runtime_config, image, image_tokens,
750+
prefill_activations, env);
751+
} // end GCPP_ZONE before we print results.
743752

744753
// No-op if the profiler is disabled. Printing now ensures that the
745754
// `PrintResults` after prefill does not include the image token part.
746755
env.ctx.profiler.PrintResults();
756+
757+
timing_info.NotifyImageTokenDone(num_tokens);
747758
}
748759

749760
// NOLINTNEXTLINE(google-readability-namespace-comments)
@@ -814,13 +825,13 @@ void Gemma::GenerateBatch(const RuntimeConfig& runtime_config,
814825

815826
void Gemma::GenerateImageTokens(const RuntimeConfig& runtime_config,
816827
size_t seq_len, const Image& image,
817-
ImageTokens& image_tokens,
818-
MatMulEnv& env) const {
828+
ImageTokens& image_tokens, MatMulEnv& env,
829+
TimingInfo& timing_info) const {
819830
env.ctx.pools.MaybeStartSpinning(runtime_config.use_spinning);
820831

821832
HWY_DYNAMIC_DISPATCH(GenerateImageTokensT)(model_.Config(), runtime_config,
822833
seq_len, weights_, image,
823-
image_tokens, env);
834+
image_tokens, env, timing_info);
824835

825836
env.ctx.pools.MaybeStopSpinning(runtime_config.use_spinning);
826837
}

gemma/gemma.h

Lines changed: 32 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -65,6 +65,21 @@ class ContinuousQBatch : public QBatch {
6565
};
6666

6767
struct TimingInfo {
68+
void NotifyImageTokenStart() { image_tokens_start = hwy::platform::Now(); }
69+
70+
void NotifyImageTokenDone(size_t tokens) {
71+
image_tokens_duration = hwy::platform::Now() - image_tokens_start;
72+
image_tokens = tokens;
73+
74+
if (verbosity >= 1) {
75+
fprintf(stderr,
76+
"\n\n[ Timing info ] Image token generation took: %d ms (%.1f "
77+
"tok/sec)\n",
78+
static_cast<int>(image_tokens_duration * 1E3),
79+
image_tokens / image_tokens_duration);
80+
}
81+
}
82+
6883
// be sure to populate prefill_start before calling NotifyPrefill.
6984
void NotifyPrefill(size_t tokens) {
7085
prefill_duration = hwy::platform::Now() - prefill_start;
@@ -87,8 +102,8 @@ struct TimingInfo {
87102
fprintf(stderr,
88103
"\n\n[ Timing info ] Prefill: %d ms for %zu prompt tokens "
89104
"(%.2f tokens / sec); Time to first token: %d ms\n",
90-
static_cast<int>(prefill_duration * 1000), prefill_tokens,
91-
prefill_tok_sec, static_cast<int>(time_to_first_token * 1000));
105+
static_cast<int>(prefill_duration * 1E3), prefill_tokens,
106+
prefill_tok_sec, static_cast<int>(time_to_first_token * 1E3));
92107
}
93108
}
94109
if (HWY_UNLIKELY(verbosity >= 2 && tokens_generated % 1024 == 0)) {
@@ -110,20 +125,27 @@ struct TimingInfo {
110125
fprintf(stderr,
111126
"\n[ Timing info ] Generate: %d ms for %zu tokens (%.2f tokens / "
112127
"sec)\n",
113-
static_cast<int>(generate_duration * 1000), tokens_generated,
128+
static_cast<int>(generate_duration * 1E3), tokens_generated,
114129
gen_tok_sec);
115130
}
116131
}
117132

118-
int verbosity = 0;
119-
double prefill_start = 0;
120-
double generate_start = 0;
121-
double prefill_duration = 0;
133+
double image_tokens_start = 0.0;
134+
double image_tokens_duration = 0.0;
135+
size_t image_tokens = 0;
136+
137+
double prefill_start = 0.0;
138+
double prefill_duration = 0.0;
122139
size_t prefill_tokens = 0;
123-
double time_to_first_token = 0;
124-
double generate_duration = 0;
140+
141+
double generate_start = 0.0;
142+
double generate_duration = 0.0;
125143
size_t tokens_generated = 0;
144+
145+
double time_to_first_token = 0.0;
126146
size_t generation_steps = 0;
147+
148+
int verbosity = 0;
127149
};
128150

129151
// After construction, all methods are const and thread-compatible if using
@@ -173,7 +195,7 @@ class Gemma {
173195
// Generates the image tokens by running the image encoder ViT.
174196
void GenerateImageTokens(const RuntimeConfig& runtime_config, size_t seq_len,
175197
const Image& image, ImageTokens& image_tokens,
176-
MatMulEnv& env) const;
198+
MatMulEnv& env, TimingInfo& timing_info) const;
177199

178200
private:
179201
BlobReader reader_;

gemma/run.cc

Lines changed: 3 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -99,6 +99,8 @@ void ReplGemma(const GemmaArgs& args, const Gemma& gemma, KVCache& kv_cache,
9999
size_t prompt_size = 0;
100100
const ModelConfig& config = gemma.Config();
101101

102+
TimingInfo timing_info = {.verbosity = inference.verbosity};
103+
102104
const bool have_image = !inference.image_file.path.empty();
103105
Image image;
104106
const size_t pool_dim = config.vit_config.pool_dim;
@@ -117,15 +119,8 @@ void ReplGemma(const GemmaArgs& args, const Gemma& gemma, KVCache& kv_cache,
117119
image.Resize(image_size, image_size);
118120
RuntimeConfig runtime_config = {.verbosity = verbosity,
119121
.use_spinning = args.threading.spin};
120-
double image_tokens_start = hwy::platform::Now();
121122
gemma.GenerateImageTokens(runtime_config, kv_cache.SeqLen(), image,
122-
image_tokens, env);
123-
if (verbosity >= 1) {
124-
double image_tokens_duration = hwy::platform::Now() - image_tokens_start;
125-
fprintf(stderr,
126-
"\n\n[ Timing info ] Image token generation took: %d ms\n",
127-
static_cast<int>(image_tokens_duration * 1000));
128-
}
123+
image_tokens, env, timing_info);
129124
}
130125

131126
// callback function invoked for each generated token.
@@ -188,7 +183,6 @@ void ReplGemma(const GemmaArgs& args, const Gemma& gemma, KVCache& kv_cache,
188183
}
189184

190185
// Set up runtime config.
191-
TimingInfo timing_info = {.verbosity = inference.verbosity};
192186
RuntimeConfig runtime_config = {.verbosity = inference.verbosity,
193187
.batch_stream_token = batch_stream_token,
194188
.use_spinning = args.threading.spin};

paligemma/paligemma_helper.cc

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -29,7 +29,8 @@ void PaliGemmaHelper::InitVit(const std::string& path) {
2929
image.Resize(image_size, image_size);
3030
RuntimeConfig runtime_config = {.verbosity = 0};
3131
gemma.GenerateImageTokens(runtime_config, env_->MutableKVCache().SeqLen(),
32-
image, *image_tokens_, env_->MutableEnv());
32+
image, *image_tokens_, env_->MutableEnv(),
33+
timing_info_);
3334
}
3435

3536
std::string PaliGemmaHelper::GemmaReply(const std::string& prompt_text) const {

paligemma/paligemma_helper.h

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,7 +3,9 @@
33

44
#include <memory>
55
#include <string>
6+
67
#include "evals/benchmark_helper.h"
8+
#include "gemma/gemma.h"
79
#include "gemma/gemma_args.h"
810

911
namespace gcpp {
@@ -18,6 +20,7 @@ class PaliGemmaHelper {
1820
private:
1921
std::unique_ptr<ImageTokens> image_tokens_;
2022
GemmaEnv* env_;
23+
TimingInfo timing_info_;
2124
};
2225

2326
} // namespace gcpp

python/gemma_py.cc

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -183,7 +183,8 @@ class GemmaModel {
183183
env_.MutableEnv().ctx.allocator, gcpp::MatPadding::kOdd));
184184
gcpp::RuntimeConfig runtime_config = {.verbosity = 0};
185185
gemma.GenerateImageTokens(runtime_config, env_.MutableKVCache().SeqLen(),
186-
c_image, *image_tokens_, env_.MutableEnv());
186+
c_image, *image_tokens_, env_.MutableEnv(),
187+
timing_info_);
187188
}
188189

189190
// Generates a response to the given prompt, using the last set image.
@@ -244,6 +245,7 @@ class GemmaModel {
244245
private:
245246
gcpp::GemmaEnv env_;
246247
std::unique_ptr<gcpp::ImageTokens> image_tokens_;
248+
gcpp::TimingInfo timing_info_;
247249
float last_prob_;
248250
};
249251

0 commit comments

Comments
 (0)