diff --git a/.github/workflows/cuda.yml b/.github/workflows/cuda.yml index 92c75413160..39054f27eac 100644 --- a/.github/workflows/cuda.yml +++ b/.github/workflows/cuda.yml @@ -480,8 +480,11 @@ jobs: pip install gguf python -m pytest examples/models/gemma4_31b/tests/ --ignore=examples/models/gemma4_31b/tests/test_mlx_pipeline.py -v -o "addopts=" - # Muse Glimmer batched CUDA export on a tiny model + # Muse Glimmer batched CUDA export on a tiny model, and the batched + # runner end to end on it python -m pytest examples/models/muse-glimmer/tests/test_cuda_batching_pipeline.py -v -o "addopts=" + make muse-glimmer-cuda + python -m pytest examples/models/muse-glimmer/tests/test_run_solo_batching.py -v -o "addopts=" unittest-cuda-runtime: name: unittest-cuda-runtime diff --git a/Makefile b/Makefile index f0fee3adfad..9e9b26f82a7 100644 --- a/Makefile +++ b/Makefile @@ -492,6 +492,7 @@ muse-glimmer-cuda: @echo "" @echo "✓ Build complete!" @echo " Solo runner: cmake-out/examples/models/muse-glimmer/solo_runner" + @echo " Batched runner: cmake-out/examples/models/muse-glimmer/run_solo_batching" @echo " DFlash runner: cmake-out/examples/models/muse-glimmer/dflash_runner" @echo " Worker: cmake-out/examples/models/muse-glimmer/muse_glimmer_worker" diff --git a/examples/models/muse-glimmer/CMakeLists.txt b/examples/models/muse-glimmer/CMakeLists.txt index ff7da644771..c14209ed17d 100644 --- a/examples/models/muse-glimmer/CMakeLists.txt +++ b/examples/models/muse-glimmer/CMakeLists.txt @@ -189,3 +189,19 @@ if(TARGET mlxdelegate) executorch_target_copy_mlx_metallib(dflash_runner) executorch_target_copy_mlx_metallib(muse_glimmer_worker) endif() + +# Batched text generation on the batching extension (CUDA only): runs what +# export/export_solo_batching.py writes. +if(EXECUTORCH_BUILD_CUDA) + add_executable(run_solo_batching runtime/runners/run_solo_batching.cpp) + target_include_directories( + run_solo_batching PUBLIC ${_common_include_directories} ${_json_include} + ) + target_link_libraries( + run_solo_batching PUBLIC ${link_libraries} cuda_batching Threads::Threads + ) + # CUDA AOTI blobs resolve symbols from the host executable. + if(NOT APPLE AND NOT MSVC) + target_link_options(run_solo_batching PRIVATE "LINKER:--export-dynamic") + endif() +endif() diff --git a/examples/models/muse-glimmer/CMakePresets.json b/examples/models/muse-glimmer/CMakePresets.json index 4cc407d4081..07f4a058dea 100644 --- a/examples/models/muse-glimmer/CMakePresets.json +++ b/examples/models/muse-glimmer/CMakePresets.json @@ -43,7 +43,7 @@ "name": "muse-glimmer-cuda", "displayName": "Build Muse Glimmer runner (CUDA)", "configurePreset": "muse-glimmer-cuda", - "targets": ["solo_runner", "dflash_runner", "muse_glimmer_worker"] + "targets": ["solo_runner", "run_solo_batching", "dflash_runner", "muse_glimmer_worker"] }, { "name": "muse-glimmer-mlx", diff --git a/examples/models/muse-glimmer/runtime/runners/run_solo_batching.cpp b/examples/models/muse-glimmer/runtime/runners/run_solo_batching.cpp new file mode 100644 index 00000000000..7fe50f8343e --- /dev/null +++ b/examples/models/muse-glimmer/runtime/runners/run_solo_batching.cpp @@ -0,0 +1,497 @@ +/* + * Copyright (c) Meta Platforms, Inc. and affiliates. + * All rights reserved. + * + * This source code is licensed under the BSD-style license found in the + * LICENSE file in the root directory of this source tree. + */ + +// Batched Muse Glimmer text generation on CUDA, on the batching extension: +// CudaExecutor runs the artifact export_solo_batching.py writes, the batching +// Runner carries every prompt's generation, and DecodeFirstScheduler packs +// their decodes and prefill chunks into shared forwards. +// +// Each prompt is one session. Generated text is printed per prompt, and +// --report_json writes what CI checks: each generation, GPU memory at each +// stage, and the KV pool's usage. + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include +#include + +DEFINE_string(model_path, "", "The .pte export_solo_batching.py wrote."); +DEFINE_string(data_path, "", "Its aoti_cuda_blob.ptd."); +DEFINE_string(tokenizer_path, "", "tokenizer.json."); +DEFINE_string( + prompt_file, + "", + "A file holding one already-formatted prompt, verbatim."); +DEFINE_string( + prompts_file, + "", + "A file holding one already-formatted prompt per line; \\n in a line is " + "a newline. Combined with --prompt_file, that prompt comes first."); +DEFINE_string( + prompt_tokens_file, + "", + "Instead of text prompts: one prompt per line as space-separated token " + "ids, BOS included. --tokenizer_path is then optional."); +DEFINE_int32(max_new_tokens, 512, "Maximum tokens generated per prompt."); +DEFINE_double(temperature, 0.0, "Sampling temperature; 0 is greedy."); +DEFINE_double(top_p, 1.0, "Nucleus sampling probability."); +DEFINE_int32(top_k, 0, "Top-k sampling limit; 0 disables it."); +DEFINE_uint64(seed, 42, "Per-generation sampling seed."); +DEFINE_int32(max_sessions, 4, "Resident sessions the KV pool reserves for."); +DEFINE_int32( + max_session_tokens, + 4096, + "Tokens one session may hold: its prompt plus what it generates."); +DEFINE_int32( + kv_initial_capacity, + -1, + "Initial KV pool rows; -1 keeps the cache's default. It grows on demand."); +DEFINE_int32( + max_decode_sequences, + 0, + "Decodes admitted to one forward; 0 admits one per session."); +DEFINE_bool(cuda_graph, true, "Capture decode into a CUDA graph."); +DEFINE_bool( + weight_sharing, + true, + "Share one copy of the weights between decode and prefill."); +DEFINE_int32(bos_id, 200000, "BOS token id."); +DEFINE_int32(eos_id, 200001, "EOS token id."); +DEFINE_string(report_json, "", "Write a JSON report here."); + +namespace batching = ::executorch::extension::llm::batching; +namespace cuda_batching = ::executorch::backends::cuda::batching; +namespace metadata = ::executorch::extension::llm; +using ::executorch::extension::Module; +using ::executorch::runtime::Error; + +namespace { + +constexpr int kBFloat16 = 15; // ScalarType::BFloat16 + +struct GpuMemory { + std::size_t used = 0; + std::size_t total = 0; +}; + +GpuMemory gpu_memory() { + std::size_t free_bytes = 0; + std::size_t total_bytes = 0; + cudaMemGetInfo(&free_bytes, &total_bytes); + return {total_bytes - free_bytes, total_bytes}; +} + +std::string read_file(const std::string& path) { + std::ifstream in(path, std::ios::binary); + if (!in) { + return {}; + } + std::ostringstream out; + out << in.rdbuf(); + return out.str(); +} + +std::string unescape_newlines(const std::string& line) { + std::string out; + for (std::size_t i = 0; i < line.size(); ++i) { + if (line[i] == '\\' && i + 1 < line.size() && line[i + 1] == 'n') { + out.push_back('\n'); + ++i; + } else { + out.push_back(line[i]); + } + } + return out; +} + +std::vector> load_prompt_tokens() { + std::vector> prompts; + std::ifstream in(FLAGS_prompt_tokens_file); + std::string line; + while (std::getline(in, line)) { + std::istringstream ids(line); + std::vector tokens; + batching::Token token; + while (ids >> token) { + tokens.push_back(token); + } + if (!tokens.empty()) { + prompts.push_back(std::move(tokens)); + } + } + return prompts; +} + +std::vector load_prompts() { + std::vector prompts; + if (!FLAGS_prompt_file.empty()) { + prompts.push_back(read_file(FLAGS_prompt_file)); + } + if (!FLAGS_prompts_file.empty()) { + std::ifstream in(FLAGS_prompts_file); + std::string line; + while (std::getline(in, line)) { + if (!line.empty()) { + prompts.push_back(unescape_newlines(line)); + } + } + } + return prompts; +} + +// Turn-ending tokens: the model's EOS plus Harmony's <|eot|> and +// <|end_of_text|>. <|eom|> continues the assistant turn, so it is not one. +std::vector stop_tokens(::tokenizers::Tokenizer& tokenizer) { + std::unordered_set ids{ + static_cast(FLAGS_eos_id)}; + for (const char* piece : {"<|eot|>", "<|end_of_text|>"}) { + if (auto id = tokenizer.piece_to_id(piece); id.ok()) { + ids.insert(static_cast(*id)); + } + } + return {ids.begin(), ids.end()}; +} + +// K and V, every layer, one token. +std::int64_t kv_bytes_per_cell(const metadata::cache::CacheGeometry& geometry) { + std::int64_t bytes = 0; + for (const auto& layer : geometry.layers) { + bytes += 2 * static_cast(layer.n_kv_heads) * layer.head_dim * + 2; // bf16 + } + return bytes; +} + +const char* reason_name(const std::optional& reason) { + if (!reason) { + return "never started"; + } + switch (*reason) { + case batching::FinishReason::StopToken: + return "stop_token"; + case batching::FinishReason::NewTokenLimit: + return "token_limit"; + case batching::FinishReason::Cancelled: + return "cancelled"; + case batching::FinishReason::Failed: + return "failed"; + } + return "unknown"; +} + +struct Job { + std::string prompt; + std::vector prompt_tokens; + std::vector generated; + std::optional session; + batching::GenerationHandle handle; + std::optional reason; + std::string error; + std::string text; +}; + +} // namespace + +int main(int argc, char** argv) { + gflags::ParseCommandLineFlags(&argc, &argv, true); + ::executorch::runtime::runtime_init(); + const bool text = FLAGS_prompt_tokens_file.empty(); + if (FLAGS_model_path.empty() || (text && FLAGS_tokenizer_path.empty())) { + std::cerr << "--model_path and --tokenizer_path are required" << std::endl; + return 1; + } + std::vector jobs; + ::tokenizers::HFTokenizer tokenizer; + const bool have_tokenizer = !FLAGS_tokenizer_path.empty(); + if (have_tokenizer && + tokenizer.load(FLAGS_tokenizer_path) != ::tokenizers::Error::Ok) { + std::cerr << "could not load " << FLAGS_tokenizer_path << std::endl; + return 1; + } + if (text) { + for (auto& prompt : load_prompts()) { + Job job{std::move(prompt)}; + auto encoded = tokenizer.encode(job.prompt, /*bos=*/0, /*eos=*/0); + if (!encoded.ok()) { + std::cerr << "could not encode a prompt" << std::endl; + return 1; + } + job.prompt_tokens.push_back(static_cast(FLAGS_bos_id)); + for (const auto token : *encoded) { + job.prompt_tokens.push_back(static_cast(token)); + } + jobs.push_back(std::move(job)); + } + } else { + for (auto& tokens : load_prompt_tokens()) { + Job job; + job.prompt_tokens = std::move(tokens); + jobs.push_back(std::move(job)); + } + } + if (jobs.empty()) { + std::cerr << "no prompts: pass --prompt_file, --prompts_file or " + "--prompt_tokens_file" + << std::endl; + return 1; + } + for (const Job& job : jobs) { + if (job.prompt_tokens.size() + FLAGS_max_new_tokens > + static_cast(FLAGS_max_session_tokens)) { + std::cerr << "a prompt plus --max_new_tokens exceeds --max_session_tokens" + << std::endl; + return 1; + } + } + const auto stops = have_tokenizer + ? stop_tokens(tokenizer) + : std::vector{static_cast(FLAGS_eos_id)}; + + // Create the CUDA context first, so its cost is not counted as the model's. + cudaFree(nullptr); + const GpuMemory before_load = gpu_memory(); + std::vector data_files; + if (!FLAGS_data_path.empty()) { + data_files.push_back(FLAGS_data_path); + } + auto module = std::make_unique( + FLAGS_model_path, + data_files, + Module::LoadMode::MmapUseMlockIgnoreErrors, + /*event_tracer=*/nullptr, + /*memory_allocator=*/nullptr, + /*temp_allocator=*/nullptr, + /*share_memory_arenas=*/false); + if (module->load() != Error::Ok) { + std::cerr << "could not load " << FLAGS_model_path << std::endl; + return 1; + } + const auto geometry = metadata::read_cache_geometry(*module); + if (!geometry.ok()) { + std::cerr << "the program publishes no KV cache geometry" << std::endl; + return 1; + } + const std::int64_t cell_bytes = kv_bytes_per_cell(*geometry); + + cuda_batching::CudaExecutorOptions options; + options.cuda_graph_for_decode = FLAGS_cuda_graph; + options.weight_sharing_across_methods = FLAGS_weight_sharing; + auto executor = cuda_batching::CudaExecutor::create( + std::move(module), + FLAGS_max_sessions, + FLAGS_max_session_tokens, + kBFloat16, + FLAGS_kv_initial_capacity, + options); + if (!executor.ok()) { + std::cerr << "CudaExecutor::create failed: 0x" << std::hex + << static_cast(executor.error()) << std::endl; + return 1; + } + + const std::size_t width = (*executor)->preferred_batch_tokens(); + const std::size_t decode_slots = FLAGS_max_decode_sequences > 0 + ? static_cast(FLAGS_max_decode_sequences) + : static_cast(FLAGS_max_sessions); + if (decode_slots >= width) { + std::cerr << "--max_decode_sequences leaves no room for prefill in a " + << width << "-token forward" << std::endl; + return 1; + } + auto scheduler = batching::DecodeFirstScheduler::create( + width, decode_slots, width - decode_slots); + batching::Runner runner(**executor, std::move(scheduler)); + + // Opening a session runs the engine's first step, which loads both methods. + for (Job& job : jobs) { + job.session = runner.open_session_async().get(); + if (!job.session) { + std::cerr << "could not open a session; raise --max_sessions" + << std::endl; + runner.shutdown(); + return 1; + } + } + const GpuMemory after_load = gpu_memory(); + + std::mutex mutex; + const auto generate_start = std::chrono::steady_clock::now(); + for (std::size_t i = 0; i < jobs.size(); ++i) { + batching::GenConfig config; + config.max_new_tokens = FLAGS_max_new_tokens; + config.sampling.temperature = static_cast(FLAGS_temperature); + config.sampling.top_p = static_cast(FLAGS_top_p); + config.sampling.top_k = FLAGS_top_k; + config.seed = FLAGS_seed; + config.stop_tokens = stops; + Job& job = jobs[i]; + job.handle = job.session->generate_async( + job.prompt_tokens, + config, + [&mutex, &job](const batching::GenerationUpdate& update) { + std::lock_guard guard(mutex); + job.generated.insert( + job.generated.end(), update.tokens.begin(), update.tokens.end()); + }); + } + for (Job& job : jobs) { + job.handle.wait(); + job.reason = job.handle.finish_reason(); + job.error = job.handle.error_message(); + } + const double generate_seconds = std::chrono::duration( + std::chrono::steady_clock::now() - generate_start) + .count(); + const GpuMemory after_generate = gpu_memory(); + const ::executorch::backends::cuda::OffGraphKVMetrics kv = + (*executor)->kv_metrics(); + for (Job& job : jobs) { + job.session.reset(); + } + runner.shutdown(); + const batching::EngineMetrics engine = runner.metrics(); + + bool ok = true; + for (std::size_t i = 0; i < jobs.size(); ++i) { + Job& job = jobs[i]; + std::uint64_t previous = job.prompt_tokens.back(); + for (const auto token : job.generated) { + if (!have_tokenizer) { + break; + } + if (auto piece = tokenizer.decode(previous, token); piece.ok()) { + job.text += *piece; + } + previous = token; + } + std::printf( + "=== [%zu] %zu prompt tokens, %zu generated, %s ===\n%s\n", + i, + job.prompt_tokens.size(), + job.generated.size(), + reason_name(job.reason), + job.text.c_str()); + if (!job.reason || *job.reason == batching::FinishReason::Failed || + *job.reason == batching::FinishReason::Cancelled) { + std::fprintf(stderr, "[%zu] failed: %s\n", i, job.error.c_str()); + ok = false; + } + } + const double mib = 1024.0 * 1024.0; + std::size_t generated_total = 0; + for (const Job& job : jobs) { + generated_total += job.generated.size(); + } + std::printf( + "Generate: %zu tokens across %zu sessions in %.2f s (%.1f tok/s)\n", + generated_total, + jobs.size(), + generate_seconds, + generate_seconds > 0 ? generated_total / generate_seconds : 0.0); + std::printf( + "Engine: %" PRIu64 " steps, %" PRIu64 " decode sessions, %" PRIu64 + " prefill sessions, %" PRIu64 " decode tokens, %" PRIu64 + " prefill tokens\n", + engine.steps, + engine.decode_sessions_total, + engine.prefill_sessions_total, + engine.decode_tokens_total, + engine.prefill_tokens_total); + std::printf( + "KV pool: %" PRId64 " rows allocated, %" PRId64 " cells in use, %.1f MiB, " + "%" PRId64 " growths\n", + kv.flat_capacity, + kv.logical_length, + kv.allocated_bytes / mib, + kv.growth_count); + std::printf( + "GPU memory: %.1f MiB before load, %.1f MiB after load, %.1f MiB after " + "generate\n", + before_load.used / mib, + after_load.used / mib, + after_generate.used / mib); + std::printf("GPU peak memory usage: %.1f MiB\n", after_generate.used / mib); + + if (!FLAGS_report_json.empty()) { + nlohmann::json report; + for (const Job& job : jobs) { + report["generations"].push_back( + {{"prompt", job.prompt}, + {"prompt_tokens", job.prompt_tokens.size()}, + {"tokens", job.generated}, + {"text", job.text}, + {"finish_reason", reason_name(job.reason)}, + {"error", job.error}}); + } + report["gpu"] = { + {"total_bytes", before_load.total}, + {"used_before_load_bytes", before_load.used}, + {"used_after_load_bytes", after_load.used}, + {"used_after_generate_bytes", after_generate.used}}; + report["kv"] = { + {"rows", kv.flat_capacity}, + {"cells_in_use", kv.logical_length}, + {"allocated_bytes", kv.allocated_bytes}, + {"growth_count", kv.growth_count}, + {"bytes_per_cell", cell_bytes}, + {"initial_capacity", FLAGS_kv_initial_capacity}}; + std::error_code size_error; + report["weights_bytes"] = FLAGS_data_path.empty() + ? 0 + : std::filesystem::file_size(FLAGS_data_path, size_error); + report["config"] = { + {"max_sessions", FLAGS_max_sessions}, + {"max_session_tokens", FLAGS_max_session_tokens}, + {"max_new_tokens", FLAGS_max_new_tokens}, + {"cuda_graph", FLAGS_cuda_graph}, + {"weight_sharing", FLAGS_weight_sharing}, + {"step_width", width}}; + report["timing"] = { + {"generate_seconds", generate_seconds}, + {"generated_tokens", generated_total}, + {"tokens_per_second", + generate_seconds > 0 ? generated_total / generate_seconds : 0.0}}; + report["engine"] = { + {"steps", engine.steps}, + {"steps_failed", engine.steps_failed}, + {"decode_sessions_total", engine.decode_sessions_total}, + {"prefill_sessions_total", engine.prefill_sessions_total}, + {"decode_tokens_total", engine.decode_tokens_total}, + {"prefill_tokens_total", engine.prefill_tokens_total}}; + std::ofstream out(FLAGS_report_json); + out << report.dump(2) << std::endl; + if (!out) { + std::cerr << "could not write " << FLAGS_report_json << std::endl; + return 1; + } + } + return ok ? 0 : 1; +} diff --git a/examples/models/muse-glimmer/tests/test_run_solo_batching.py b/examples/models/muse-glimmer/tests/test_run_solo_batching.py new file mode 100644 index 00000000000..1250c80b7d3 --- /dev/null +++ b/examples/models/muse-glimmer/tests/test_run_solo_batching.py @@ -0,0 +1,161 @@ +# Copyright (c) Meta Platforms, Inc. and affiliates. +# All rights reserved. +# +# This source code is licensed under the BSD-style license found in the +# LICENSE file in the root directory of this source tree. + +"""End to end: export_solo_batching.py's artifact run by run_solo_batching. + +Exports the tiny model once, then drives the built runner binary with token-id +prompts and checks its JSON report. The tiny model's weights are random, so +the checks are ones that hold whatever it generates: every generation +completes, the captured decode graph generates what eager decode does, +forwards are shared across sessions, weights load once, and the KV pool grows +with use. Tokens are not compared across batch compositions: a prompt +prefilled in a wider forward runs GEMMs of another shape, and bf16 rounding +can steer greedy decoding elsewhere. + +Needs CUDA and the runner (`make muse-glimmer-cuda`); its path may be set with +MUSE_GLIMMER_BATCHING_RUNNER. + + python -m pytest examples/models/muse-glimmer/tests/test_run_solo_batching.py -v +""" + +import json +import os +import shutil +import subprocess +import tempfile +import unittest + +import executorch.backends.cuda.quantize_op_dispatch as _quantize_op_dispatch # noqa: F401 +import torch +from executorch.examples.models.muse_glimmer.export.export_solo import ( + load_prequantized_model, +) +from executorch.examples.models.muse_glimmer.export.export_solo_batching import ( + export_batching, +) +from executorch.examples.models.muse_glimmer.tests.test_pipeline import ( + save_checkpoint, + TINY_CONFIG, +) + +_REPO = os.path.abspath(os.path.join(os.path.dirname(__file__), "..", "..", "..", "..")) +_RUNNER = os.environ.get( + "MUSE_GLIMMER_BATCHING_RUNNER", + os.path.join(_REPO, "cmake-out", "examples", "models", "muse-glimmer", "run_solo_batching"), +) +MAX_STEP = 16 +MAX_CELLS = 256 +NEW_TOKENS = 8 +MIB = 1024 * 1024 + + +class RunSoloBatchingTest(unittest.TestCase): + @classmethod + def setUpClass(cls) -> None: + if not torch.cuda.is_available(): + raise unittest.SkipTest("CUDA required") + if not os.path.exists(_RUNNER): + raise unittest.SkipTest(f"runner not built: {_RUNNER}") + cls.tmp = tempfile.mkdtemp() + ckpt = os.path.join(cls.tmp, "ckpt") + cls.artifact = os.path.join(cls.tmp, "artifact") + save_checkpoint(ckpt) + model, config = load_prequantized_model(ckpt, max_seq_len=TINY_CONFIG.max_seq_len) + export_batching( + model, config, cls.artifact, max_step_tokens=MAX_STEP, max_cells=MAX_CELLS + ) + generator = torch.Generator().manual_seed(0) + vocab = TINY_CONFIG.vocab_size + + def prompt(length): + return torch.randint(0, vocab, (length,), generator=generator).tolist() + + # One token, a prefill that runs as decodes (< 5), one prefill, and one + # wider than a forward so it slices. + cls.prompts = [prompt(1), prompt(3), prompt(9), prompt(30)] + + @classmethod + def tearDownClass(cls) -> None: + shutil.rmtree(getattr(cls, "tmp", ""), ignore_errors=True) + + def run_prompts(self, prompts, **flags) -> dict: + prompts_file = os.path.join(self.tmp, "prompts.txt") + report_file = os.path.join(self.tmp, "report.json") + with open(prompts_file, "w") as f: + for tokens in prompts: + f.write(" ".join(map(str, tokens)) + "\n") + args = { + "model_path": os.path.join(self.artifact, "model.pte"), + "data_path": os.path.join(self.artifact, "aoti_cuda_blob.ptd"), + "prompt_tokens_file": prompts_file, + "max_new_tokens": NEW_TOKENS, + "max_sessions": 4, + "max_session_tokens": TINY_CONFIG.max_seq_len, + "kv_initial_capacity": 16, + # Never a stop token: every generation runs to the limit. + "eos_id": TINY_CONFIG.vocab_size + 1, + "report_json": report_file, + } + args.update(flags) + command = [_RUNNER] + [ + f"--{key}={str(value).lower() if isinstance(value, bool) else value}" + for key, value in args.items() + ] + result = subprocess.run(command, capture_output=True, text=True, timeout=600) + self.assertEqual(result.returncode, 0, result.stderr[-4000:]) + with open(report_file) as f: + return json.load(f) + + def test_concurrent_generations_complete_in_shared_forwards(self) -> None: + duplicate = self.prompts[2] + report = self.run_prompts(self.prompts + [duplicate], max_sessions=5, max_session_tokens=48) + generations = report["generations"] + self.assertEqual(len(generations), 5) + for generation in generations: + self.assertEqual(generation["finish_reason"], "token_limit") + self.assertEqual(len(generation["tokens"]), NEW_TOKENS) + self.assertTrue( + all(0 <= t < TINY_CONFIG.vocab_size for t in generation["tokens"]) + ) + # Run one after another, every generated token would take a forward of + # its own; sharing forwards takes far fewer. + engine = report["engine"] + self.assertEqual(engine["steps_failed"], 0) + self.assertEqual(engine["decode_tokens_total"], 5 * (NEW_TOKENS - 1)) + self.assertLess(engine["steps"], engine["decode_tokens_total"]) + + def test_captured_decode_graph_generates_what_eager_decode_does(self) -> None: + prompt = [self.prompts[3]] + graph = self.run_prompts(prompt, cuda_graph=True) + eager = self.run_prompts(prompt, cuda_graph=False) + self.assertEqual( + graph["generations"][0]["tokens"], eager["generations"][0]["tokens"] + ) + + def test_weights_load_once_whatever_the_session_count(self) -> None: + one = self.run_prompts(self.prompts[:1], max_sessions=1) + many = self.run_prompts(self.prompts[:1], max_sessions=4, max_session_tokens=64) + load = [ + r["gpu"]["used_after_load_bytes"] - r["gpu"]["used_before_load_bytes"] + for r in (one, many) + ] + # The pool is allocated at the first step, not at load, and grows with + # use: reserving for more sessions costs nothing up front. + self.assertLess(abs(load[0] - load[1]), 64 * MIB) + # Loading is the weights once, not once per method, plus a fixed cost. + self.assertLess(load[0], 2 * one["weights_bytes"] + 256 * MIB) + + def test_kv_pool_grows_with_use(self) -> None: + report = self.run_prompts(self.prompts, max_session_tokens=48) + kv = report["kv"] + tokens = sum(len(p) + NEW_TOKENS for p in self.prompts) + self.assertGreaterEqual(kv["rows"], kv["cells_in_use"]) + self.assertGreaterEqual(kv["growth_count"], 1) + # Geometric growth past a 16-row start: at most twice what was needed. + self.assertLessEqual(kv["rows"], max(16, 2 * (tokens + MAX_STEP))) + self.assertEqual(kv["allocated_bytes"], kv["rows"] * kv["bytes_per_cell"]) + # Far short of reserving every session's full context up front. + self.assertLess(kv["rows"], 4 * 48)