1- // C++ DFlash runner for benchmarking the exported Gemma4-31B target and draft models.
1+ // C++ DFlash runner for benchmarking the exported Gemma4-31B target and draft
2+ // models.
23
34#include < gflags/gflags.h>
45
1011#include < algorithm>
1112#include < chrono>
1213#include < cstdint>
13- #include < cstring>
1414#include < cstdio>
15+ #include < cstring>
1516#include < set>
1617#include < string>
1718#include < vector>
1819
19- using executorch::extension::Module;
2020using executorch::extension::make_tensor_ptr;
21+ using executorch::extension::Module;
2122using executorch::extension::TensorPtr;
22- using executorch::runtime::EValue;
2323using executorch::runtime::Error;
24+ using executorch::runtime::EValue;
2425
25- DEFINE_string (target_pte, " gemma4_31b_dflash_exports_mlx/model.pte" , " Target .pte path." );
26+ DEFINE_string (
27+ target_pte,
28+ " gemma4_31b_dflash_exports_mlx/model.pte" ,
29+ " Target .pte path." );
2630DEFINE_string (draft_pte, " gemma4_31b_dflash_draft.pte" , " Draft .pte path." );
27- DEFINE_string (tokenizer_path, " gemma-4-31B-it-HQQ-INT4/tokenizer.json" , " Tokenizer path." );
31+ DEFINE_string (
32+ tokenizer_path,
33+ " gemma-4-31B-it-HQQ-INT4/tokenizer.json" ,
34+ " Tokenizer path." );
2835DEFINE_string (prompt, " The capital of France is" , " Prompt text." );
2936DEFINE_int32 (max_new_tokens, 64 , " Maximum tokens to generate." );
3037DEFINE_int32 (block_size, 16 , " DFlash draft block size." );
31- DEFINE_int32 (mask_id, 4 , " DFlash mask token id (from draft checkpoint config)." );
38+ DEFINE_int32 (
39+ mask_id,
40+ 4 ,
41+ " DFlash mask token id (from draft checkpoint config)." );
3242DEFINE_bool (raw_prompt, false , " Skip chat-template wrapping." );
3343DEFINE_bool (verbose, false , " Print per-round timing/acceptance debug output." );
3444
@@ -72,7 +82,8 @@ int64_t first_mismatch(
7282 return static_cast <int64_t >(draft_ids.size ());
7383}
7484
75- // Convert BF16 hidden states returned by the target model to FP32 for the draft model.
85+ // Convert BF16 hidden states returned by the target model to FP32 for the draft
86+ // model.
7687std::vector<float > bf16_tensor_to_fp32 (const executorch::aten::Tensor& t) {
7788 int64_t numel = t.numel ();
7889 const uint16_t * src = reinterpret_cast <const uint16_t *>(t.const_data_ptr ());
@@ -84,7 +95,7 @@ std::vector<float> bf16_tensor_to_fp32(const executorch::aten::Tensor& t) {
8495 return out;
8596}
8697
87- }
98+ } // namespace
8899
89100int main (int argc, char ** argv) {
90101 gflags::ParseCommandLineFlags (&argc, &argv, true );
@@ -122,15 +133,20 @@ int main(int argc, char** argv) {
122133 printf (" Prompt tokens: %lld\n " , (long long )prompt_len);
123134
124135 std::vector<int64_t > input_pos_vec (prompt_len);
125- for (int64_t i = 0 ; i < prompt_len; ++i) input_pos_vec[i] = i;
136+ for (int64_t i = 0 ; i < prompt_len; ++i)
137+ input_pos_vec[i] = i;
126138
127- auto prompt_ids_tensor =
128- make_tensor_ptr ({1 , (int )prompt_len}, prompt_ids.data (), executorch::aten::ScalarType::Long);
139+ auto prompt_ids_tensor = make_tensor_ptr (
140+ {1 , (int )prompt_len},
141+ prompt_ids.data (),
142+ executorch::aten::ScalarType::Long);
129143 auto input_pos_tensor = make_tensor_ptr (
130- {(int )prompt_len}, input_pos_vec.data (), executorch::aten::ScalarType::Long);
144+ {(int )prompt_len},
145+ input_pos_vec.data (),
146+ executorch::aten::ScalarType::Long);
131147
132- auto prefill_result =
133- target_module. forward ( {EValue (prompt_ids_tensor), EValue (input_pos_tensor)});
148+ auto prefill_result = target_module. forward (
149+ {EValue (prompt_ids_tensor), EValue (input_pos_tensor)});
134150 if (!prefill_result.ok ()) {
135151 ET_LOG (Error, " Prefill forward failed" );
136152 return 1 ;
@@ -148,31 +164,36 @@ int main(int argc, char** argv) {
148164 int64_t accepted_total = 0 ;
149165 int64_t emitted_total = 0 ;
150166
151-
152-
153167 int64_t hidden_concat_dim = hidden.sizes ()[2 ];
154168 std::vector<float > hidden_history = bf16_tensor_to_fp32 (hidden);
155169 int64_t hidden_len = prompt_len;
156170
157- // Warm up the draft model to avoid including one-time MLX JIT compilation in benchmark timings.
171+ // Warm up the draft model to avoid including one-time MLX JIT compilation in
172+ // benchmark timings.
158173 {
159174 std::vector<int64_t > warm_input_vec (FLAGS_block_size);
160175 warm_input_vec[0 ] = last_token;
161- for (int64_t i = 1 ; i < FLAGS_block_size; ++i) warm_input_vec[i] = FLAGS_mask_id;
176+ for (int64_t i = 1 ; i < FLAGS_block_size; ++i)
177+ warm_input_vec[i] = FLAGS_mask_id;
162178 auto warm_input_tensor = make_tensor_ptr (
163- {1 , (int )FLAGS_block_size}, warm_input_vec.data (), executorch::aten::ScalarType::Long);
179+ {1 , (int )FLAGS_block_size},
180+ warm_input_vec.data (),
181+ executorch::aten::ScalarType::Long);
164182 auto warm_hidden_tensor = make_tensor_ptr (
165183 {1 , (int )hidden_len, (int )hidden_concat_dim},
166184 hidden_history.data (),
167185 executorch::aten::ScalarType::Float);
168186 std::vector<int64_t > warm_pos_vec (hidden_len + FLAGS_block_size);
169- for (int64_t i = 0 ; i < hidden_len + FLAGS_block_size; ++i) warm_pos_vec[i] = i;
187+ for (int64_t i = 0 ; i < hidden_len + FLAGS_block_size; ++i)
188+ warm_pos_vec[i] = i;
170189 auto warm_pos_tensor = make_tensor_ptr (
171190 {1 , (int )(hidden_len + FLAGS_block_size)},
172191 warm_pos_vec.data (),
173192 executorch::aten::ScalarType::Long);
174193 auto warm_result = draft_module.forward (
175- {EValue (warm_input_tensor), EValue (warm_hidden_tensor), EValue (warm_pos_tensor)});
194+ {EValue (warm_input_tensor),
195+ EValue (warm_hidden_tensor),
196+ EValue (warm_pos_tensor)});
176197 if (!warm_result.ok ()) {
177198 ET_LOG (Error, " Draft warm-up forward failed" );
178199 return 1 ;
@@ -188,23 +209,31 @@ int main(int argc, char** argv) {
188209
189210 std::vector<int64_t > draft_input_vec (bs);
190211 draft_input_vec[0 ] = last_token;
191- for (int64_t i = 1 ; i < bs; ++i) draft_input_vec[i] = FLAGS_mask_id;
192- auto draft_input_tensor =
193- make_tensor_ptr ({1 , (int )bs}, draft_input_vec.data (), executorch::aten::ScalarType::Long);
212+ for (int64_t i = 1 ; i < bs; ++i)
213+ draft_input_vec[i] = FLAGS_mask_id;
214+ auto draft_input_tensor = make_tensor_ptr (
215+ {1 , (int )bs},
216+ draft_input_vec.data (),
217+ executorch::aten::ScalarType::Long);
194218
195219 auto hidden_tensor = make_tensor_ptr (
196220 {1 , (int )hidden_len, (int )hidden_concat_dim},
197221 hidden_history.data (),
198222 executorch::aten::ScalarType::Float);
199223
200224 std::vector<int64_t > draft_pos_vec (hidden_len + bs);
201- for (int64_t i = 0 ; i < hidden_len + bs; ++i) draft_pos_vec[i] = i;
225+ for (int64_t i = 0 ; i < hidden_len + bs; ++i)
226+ draft_pos_vec[i] = i;
202227 auto draft_pos_tensor = make_tensor_ptr (
203- {1 , (int )(hidden_len + bs)}, draft_pos_vec.data (), executorch::aten::ScalarType::Long);
228+ {1 , (int )(hidden_len + bs)},
229+ draft_pos_vec.data (),
230+ executorch::aten::ScalarType::Long);
204231
205232 double dt0 = now_ms ();
206233 auto draft_result = draft_module.forward (
207- {EValue (draft_input_tensor), EValue (hidden_tensor), EValue (draft_pos_tensor)});
234+ {EValue (draft_input_tensor),
235+ EValue (hidden_tensor),
236+ EValue (draft_pos_tensor)});
208237 double draft_exec_ms = now_ms () - dt0;
209238 if (!draft_result.ok ()) {
210239 ET_LOG (Error, " Draft forward failed at round %lld" , (long long )rounds);
@@ -215,15 +244,21 @@ int main(int argc, char** argv) {
215244
216245 std::vector<int64_t > verify_input_vec;
217246 verify_input_vec.push_back (last_token);
218- verify_input_vec.insert (verify_input_vec.end (), draft_ids.begin (), draft_ids.end ());
247+ verify_input_vec.insert (
248+ verify_input_vec.end (), draft_ids.begin (), draft_ids.end ());
219249 int64_t verify_len = static_cast <int64_t >(verify_input_vec.size ());
220250 auto verify_input_tensor = make_tensor_ptr (
221- {1 , (int )verify_len}, verify_input_vec.data (), executorch::aten::ScalarType::Long);
251+ {1 , (int )verify_len},
252+ verify_input_vec.data (),
253+ executorch::aten::ScalarType::Long);
222254
223255 std::vector<int64_t > verify_pos_vec (verify_len);
224- for (int64_t i = 0 ; i < verify_len; ++i) verify_pos_vec[i] = pos + i;
256+ for (int64_t i = 0 ; i < verify_len; ++i)
257+ verify_pos_vec[i] = pos + i;
225258 auto verify_pos_tensor = make_tensor_ptr (
226- {(int )verify_len}, verify_pos_vec.data (), executorch::aten::ScalarType::Long);
259+ {(int )verify_len},
260+ verify_pos_vec.data (),
261+ executorch::aten::ScalarType::Long);
227262
228263 double vt0 = now_ms ();
229264 auto verify_result = target_module.forward (
@@ -247,14 +282,16 @@ int main(int argc, char** argv) {
247282 (long long )hidden_len);
248283 }
249284
250- std::vector<int64_t > new_tokens (draft_ids.begin (), draft_ids.begin () + accepted);
285+ std::vector<int64_t > new_tokens (
286+ draft_ids.begin (), draft_ids.begin () + accepted);
251287 new_tokens.push_back (target_ids[accepted]);
252288
253289 bool hit_eos = false ;
254290 for (size_t i = 0 ; i < new_tokens.size (); ++i) {
255291 if (kEosIds .count (new_tokens[i])) {
256292 new_tokens.resize (i + 1 );
257- accepted = std::min (accepted, static_cast <int64_t >(new_tokens.size ()) - 1 );
293+ accepted =
294+ std::min (accepted, static_cast <int64_t >(new_tokens.size ()) - 1 );
258295 hit_eos = true ;
259296 break ;
260297 }
@@ -274,7 +311,8 @@ int main(int argc, char** argv) {
274311 new_hidden_fp32.begin () + append_len * hidden_concat_dim);
275312 hidden_len += append_len;
276313
277- if (hit_eos) break ;
314+ if (hit_eos)
315+ break ;
278316 }
279317
280318 double total_ms = now_ms () - t0;
0 commit comments