Handles warmup, running ar benchmark, and printing results.
(config, engine_generate, params, decode_state, global_batch_size, cache_size, model_size, iters)
| 171 | |
| 172 | |
| 173 | def ar_benchmark(config, engine_generate, params, decode_state, global_batch_size, cache_size, model_size, iters): |
| 174 | """Handles warmup, running ar benchmark, and printing results.""" |
| 175 | rng = jax.random.PRNGKey(1234) |
| 176 | for _ in range(_WARMUP_ITERS): |
| 177 | rng, rng_generate = jax.random.split(rng) |
| 178 | decode_state, _ = engine_generate(params, decode_state, rng_generate) |
| 179 | jax.block_until_ready(decode_state) |
| 180 | |
| 181 | time_in_s, decode_state = ar_benchmark_loop( |
| 182 | config, engine_generate, params, decode_state, iters, profile_name="autoregress" |
| 183 | ) |
| 184 | seconds_per_step = time_in_s / iters |
| 185 | ar_average_ms = seconds_per_step * 1000 |
| 186 | total_throughput = global_batch_size / seconds_per_step |
| 187 | |
| 188 | GB_per_step_per_device = (model_size + cache_size) / 1e9 / jax.device_count() |
| 189 | bw_per_device = GB_per_step_per_device / seconds_per_step |
| 190 | print( |
| 191 | f"AutoRegressive results:\n" |
| 192 | f"\tAR step average time: {ar_average_ms:.3f} ms\n" |
| 193 | f"\tAR step average time per seq: {ar_average_ms/global_batch_size:.3f} ms\n" |
| 194 | f"\tAR global batch size: {global_batch_size}\n" |
| 195 | f"\tAR throughput: {total_throughput:.3f} tokens/second\n" |
| 196 | f"\tAR memory bandwidth per device: {bw_per_device:.3f} GB/s\n\n\n" |
| 197 | ) |
| 198 | |
| 199 | result_dict = { |
| 200 | "step_in_ms": ar_average_ms, |
| 201 | "step_in_ms_per_seq": ar_average_ms / global_batch_size, |
| 202 | "global_batch_size": global_batch_size, |
| 203 | "total_throughput_tokens_per_second": total_throughput, |
| 204 | "bw_per_device_GB_per_second": bw_per_device, |
| 205 | } |
| 206 | return result_dict, decode_state |
| 207 | |
| 208 | |
| 209 | def collate_results(config, results, model_size, cache_size, num_model_params, incl_config=False): |
no test coverage detected