diff --git a/.github/workflows/cuda.yml b/.github/workflows/cuda.yml index 1472146dbee..92c75413160 100644 --- a/.github/workflows/cuda.yml +++ b/.github/workflows/cuda.yml @@ -480,6 +480,9 @@ 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 + python -m pytest examples/models/muse-glimmer/tests/test_cuda_batching_pipeline.py -v -o "addopts=" + unittest-cuda-runtime: name: unittest-cuda-runtime needs: [changed-files, run-decision] diff --git a/examples/models/muse-glimmer/export/export_solo_batching.py b/examples/models/muse-glimmer/export/export_solo_batching.py new file mode 100644 index 00000000000..7b1a1471e60 --- /dev/null +++ b/examples/models/muse-glimmer/export/export_solo_batching.py @@ -0,0 +1,306 @@ +# 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. + +"""Export a Muse Glimmer text model for batched serving on CUDA. + +The artifact is what ``CudaExecutor`` (backends/cuda/batching) drives: two +methods over one set of weights, both taking ``(tokens[1, T], input_pos[T], +logits_to_keep[K])`` and returning float32 ``logits[K, vocab]`` for the +selected rows. + +- ``decode``: T = K = 1, static. The runtime captures it into a CUDA graph. +- ``prefill``: T and K in [5, W], dynamic. Narrower slices run as decodes. + +The KV cache lives off-graph in the cell layout: every sequence of a batch +shares one pool of ``max_cells`` per-token cells, and the runtime cache writes +each step's placement and visibility before the forward. Sampling happens on +the host, per session. + + python -m executorch.examples.models.muse_glimmer.export.export_solo_batching \\ + --gguf Muse-Glimmer-30B-KQuant-17GB-Q4_K_M.gguf --output-dir out/ +""" + +import argparse +import gc +import json +import types + +import torch +from executorch.examples.models.muse_glimmer.export import common +from executorch.examples.models.muse_glimmer.export.export_solo import ( + _solo_constant_methods, + load_and_quantize, + load_prequantized_model, +) +from executorch.examples.models.muse_glimmer.model.model import ( + MuseGlimmerConfig, + MuseGlimmerModel, +) + +# Read by CudaExecutor: the cell pool's size, which fixes its step buffers. +MAX_CELLS_METHOD = "get_offgraph_kv_max_cells" +# The fewest tokens, and selected rows, prefill is exported for; CudaExecutor +# reads it from get_min_prefill_chunk and runs narrower slices as decodes. The +# quantized linears switch kernels at M <= 4 (backends/cuda/quantize_op_dispatch), +# so a dynamic width only traces from 5 -- as in the single-sequence export. +MIN_PREFILL_TOKENS = 5 + + +def step_width_spec(method: str) -> bytes: + """``delegate_input:dim`` of input_pos[T], the step's token count. + + It names the delegate's own input order, which partitioning decides: the + static method takes (tokens, input_pos, logits_to_keep), while the dynamic + one reads logits_to_keep's symbolic size first and takes (tokens, + logits_to_keep, input_pos). The CUDA backend's + CheckOffGraphKVStepWidthPass fails the export if this is wrong. + """ + return b"2:0" if method == "prefill" else b"1:0" + + +def step_forward( + model: MuseGlimmerModel, + tokens: torch.Tensor, + input_pos: torch.Tensor, + logits_to_keep: torch.Tensor, +) -> torch.Tensor: + """Logits for the selected tokens of one packed step: [K, vocab] float32. + + The whole step runs through the decoder so every token's K/V reaches the + cache; only the rows the batch samples from go through the LM head. + """ + x = model._run_blocks(model.embed_text(tokens), input_pos) + x = model.output_norm(x[:, logits_to_keep, :]) + return model._soft_cap(model.lm_head(x))[0] + + +def cell_manifest(sequence_manifest: str, max_cells: int) -> str: + """The cell-layout lowering manifest for the model's off-graph layers.""" + manifest = json.loads(sequence_manifest) + manifest["layout"] = "cell" + manifest["max_cells"] = max_cells + return json.dumps(manifest, sort_keys=True, separators=(",", ":")) + + +def export_batching( + model: MuseGlimmerModel, + config: MuseGlimmerConfig, + output_dir: str, + max_step_tokens: int, + max_cells: int, +): + """Exports ``decode`` and ``prefill`` and writes the artifact. + + Returns the ExecutorchProgramManager so callers can inspect it. + """ + import executorch.backends.cuda.quantize_op_dispatch # noqa: F401 + import torch._inductor.config as inductor_config + from executorch.backends.cuda.cuda_backend import CudaBackend + from executorch.backends.cuda.cuda_partitioner import CudaPartitioner + from executorch.backends.cuda.passes.lower_offgraph_kv import ( + OFFGRAPH_KV_COMPILE_SPEC, + OFFGRAPH_KV_STEP_WIDTH_COMPILE_SPEC, + ) + from executorch.examples.models.muse_glimmer.source_transformations.cuda import ( + enable_offgraph_kv_cache, + offgraph_kv_cache_geometry, + ) + from executorch.exir import ( + EdgeCompileConfig, + ExecutorchBackendConfig, + to_edge_transform_and_lower, + ) + from executorch.exir.backend.compile_spec_schema import CompileSpec + from executorch.exir.passes import MemoryPlanningPass + from executorch.extension.llm.export.model_metadata import ( + write_logits_to_keep_mode, + write_max_context_len, + ) + from torch.export import Dim, export + + if not MIN_PREFILL_TOKENS <= max_step_tokens <= max_cells: + raise ValueError( + f"max_step_tokens must be in [{MIN_PREFILL_TOKENS}, max_cells], " + f"got {max_step_tokens}" + ) + if max_step_tokens >= config.max_seq_len: + raise ValueError("max_step_tokens must be below the model's context") + + inductor_config.coordinate_descent_tuning = False + inductor_config.aot_inductor.compile_wrapper_opt_level = "O0" + # The PCH path shells out to `openssl sha512` and is flaky; it is only a + # compile-time optimization. + inductor_config.aot_inductor.precompile_headers = False + if hasattr(inductor_config, "cpp_cache_precompile_headers"): + inductor_config.cpp_cache_precompile_headers = False + + manifest = cell_manifest(enable_offgraph_kv_cache(model, max_step_tokens), max_cells) + geometry = offgraph_kv_cache_geometry(model) + + width = Dim("width", min=MIN_PREFILL_TOKENS, max=max_step_tokens) + rows = Dim("rows", min=MIN_PREFILL_TOKENS, max=max_step_tokens) + programs = {} + with common.BoundMethodForward( + model, types.MethodType(step_forward, model) + ), torch.no_grad(): + print("Exporting decode (T=1)...") + programs["decode"] = export( + model, + ( + torch.zeros((1, 1), dtype=torch.long), + torch.zeros(1, dtype=torch.long), + torch.zeros(1, dtype=torch.long), + ), + strict=True, + ) + print(f"Exporting prefill (T in [{MIN_PREFILL_TOKENS}, {max_step_tokens}])...") + programs["prefill"] = export( + model, + ( + torch.zeros((1, max_step_tokens), dtype=torch.long), + torch.arange(max_step_tokens, dtype=torch.long), + torch.zeros(max_step_tokens, dtype=torch.long), + ), + dynamic_shapes=({1: width}, {0: width}, {0: rows}), + strict=True, + ) + del model + gc.collect() + torch.cuda.empty_cache() + + def partitioner(name: str) -> CudaPartitioner: + return CudaPartitioner( + [ + CudaBackend.generate_method_name_compile_spec(name), + CompileSpec("low_memory_mode", b"ON"), + CompileSpec(OFFGRAPH_KV_COMPILE_SPEC, manifest.encode()), + CompileSpec(OFFGRAPH_KV_STEP_WIDTH_COMPILE_SPEC, step_width_spec(name)), + CompileSpec("autotune_at_compile_time", b"OFF"), + ] + ) + + constant_methods = _solo_constant_methods( + config=config, + max_prefill=max_step_tokens, + activation_dtype=torch.bfloat16, + mutable_buffer_metadata=None, + has_vision=False, + max_vision_patches=0, + ) + constant_methods.update(geometry) + constant_methods.update(write_max_context_len(config.max_seq_len)) + constant_methods.update(write_logits_to_keep_mode("selected")) + constant_methods["get_min_prefill_chunk"] = MIN_PREFILL_TOKENS + constant_methods["use_sampling"] = False + constant_methods[MAX_CELLS_METHOD] = max_cells + + print("Lowering decode and prefill to ExecuTorch (CUDA)...") + edge = to_edge_transform_and_lower( + programs, + partitioner={name: [partitioner(name)] for name in programs}, + compile_config=EdgeCompileConfig( + _check_ir_validity=False, + _skip_dim_order=True, + ), + constant_methods=constant_methods, + ) + del programs + gc.collect() + torch.cuda.empty_cache() + + # Logits come back to the host: each session samples its own rows there. + et_program = edge.to_executorch( + config=ExecutorchBackendConfig( + extract_delegate_segments=True, + do_quant_fusion_and_const_prop=True, + memory_planning_pass=MemoryPlanningPass(alloc_graph_input=False), + emit_mutable_buffer_names=True, + ), + ) + del edge + gc.collect() + + common.save_pte(et_program, output_dir, None) + if torch.cuda.is_available(): + peak_mb = torch.cuda.max_memory_allocated() / (1024**2) + print(f"EXPORT_GPU_PEAK_MEMORY_MB: {peak_mb:.1f}") + print("Done.") + return et_program + + +def main() -> None: + parser = argparse.ArgumentParser( + description="Export Muse Glimmer for batched serving on CUDA." + ) + src = parser.add_mutually_exclusive_group(required=True) + src.add_argument("--gguf", default=None, help="Path to a GGUF checkpoint.") + src.add_argument( + "--prequantized", default=None, help="Path to a quantized checkpoint dir." + ) + src.add_argument( + "--checkpoint-dir", + default=None, + help="Path to a consolidated bf16 checkpoint; quantized with --quant-recipe.", + ) + parser.add_argument("--quant-recipe", default="default") + parser.add_argument("--output-dir", default="./muse_glimmer_batching") + parser.add_argument( + "--max-seq-len", + type=int, + default=131072, + help="Model context: the most tokens one session may hold.", + ) + parser.add_argument( + "--max-step-tokens", + type=int, + default=512, + help="prefill's widest step, W: the most tokens one forward carries.", + ) + parser.add_argument( + "--max-cells", + type=int, + default=65536, + help="Cells in the shared KV pool: the tokens every resident session " + "together may hold. Memory grows with use, up to this.", + ) + args = parser.parse_args() + if not torch.cuda.is_available(): + parser.error("CUDA is required.") + + if args.gguf: + from executorch.examples.models.muse_glimmer.loaders.checkpoint_loader import ( + load_gguf_model, + ) + + model, config = load_gguf_model( + args.gguf, + max_seq_len=args.max_seq_len, + backend="cuda", + activation_dtype=torch.bfloat16, + ) + elif args.prequantized: + model, config = load_prequantized_model( + args.prequantized, max_seq_len=args.max_seq_len, backend="cuda" + ) + else: + model, config = load_and_quantize( + args.checkpoint_dir, + args.quant_recipe, + max_seq_len=args.max_seq_len, + backend="cuda", + ) + + export_batching( + model, + config, + args.output_dir, + max_step_tokens=args.max_step_tokens, + max_cells=args.max_cells, + ) + + +if __name__ == "__main__": + main() diff --git a/examples/models/muse-glimmer/tests/test_cuda_batching_pipeline.py b/examples/models/muse-glimmer/tests/test_cuda_batching_pipeline.py new file mode 100644 index 00000000000..c882b336d32 --- /dev/null +++ b/examples/models/muse-glimmer/tests/test_cuda_batching_pipeline.py @@ -0,0 +1,263 @@ +# 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. + +"""Tests for the batched CUDA export (export_solo_batching.py) on a tiny model. + +The eager checks run the off-graph model against the neutral cache's +references; the export check lowers ``decode`` and ``prefill`` and inspects the +artifact. Requires CUDA for the export. + + python -m pytest examples/models/muse-glimmer/tests/test_cuda_batching_pipeline.py -v +""" + +import copy +import json +import os +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 ( + cell_manifest, + export_batching, + MAX_CELLS_METHOD, + step_forward, +) +from executorch.examples.models.muse_glimmer.source_transformations.cuda import ( + enable_offgraph_kv_cache, +) +from executorch.examples.models.muse_glimmer.tests.test_pipeline import ( + build_random_tiny_model, + save_checkpoint, + TINY_CONFIG, +) +from executorch.extension.llm.cache.reference_cache import ( + CacheConfig, + CellReferenceCache, + LayerPolicy, + SequenceReferenceCache, +) +from executorch.extension.llm.cache.update_and_attend import REGISTRY + + +def _cache_config(model, capacity: int) -> CacheConfig: + attn = model.layers[0].self_attn + return CacheConfig( + n_layers=TINY_CONFIG.n_layers, + n_kv_heads=attn.n_kv_heads, + head_dim=attn.head_dim, + capacity=capacity, + layers=tuple( + ( + LayerPolicy.ring(layer.self_attn.window_size) + if layer.self_attn.is_sliding + else LayerPolicy.flat() + ) + for layer in model.layers + ), + ) + + +class BatchingStepForwardTest(unittest.TestCase): + """The exported step computes what the model does, per packed sequence.""" + + def setUp(self) -> None: + self.model = build_random_tiny_model() + self.off_graph = copy.deepcopy(self.model) + enable_offgraph_kv_cache(self.off_graph, 16) + + def _install(self, cache) -> str: + key = f"muse-batching-{id(self)}-{id(cache)}" + REGISTRY.install(key, cache) + self.addCleanup(REGISTRY.uninstall, key) + return key + + def test_selected_rows_match_the_full_forward(self) -> None: + key = self._install( + SequenceReferenceCache(_cache_config(self.off_graph, TINY_CONFIG.max_seq_len)) + ) + generator = torch.Generator().manual_seed(0) + start = 0 + for length, keep in ((12, [3, 11]), (1, [0]), (6, [5])): + tokens = torch.randint( + 0, TINY_CONFIG.vocab_size, (1, length), generator=generator + ) + input_pos = torch.arange(start, start + length) + with torch.no_grad(): + expected = self.model(tokens, input_pos)[0, keep] + with REGISTRY.active(key): + actual = step_forward( + self.off_graph, tokens, input_pos, torch.tensor(keep) + ) + self.assertEqual(actual.dtype, torch.float32) + self.assertEqual(tuple(actual.shape), (len(keep), TINY_CONFIG.vocab_size)) + self.assertLess((expected - actual).abs().max().item(), 5e-2) + start += length + + def test_packed_sequences_match_each_alone(self) -> None: + # Two prompts prefilled in one step, then decoded together, over the + # neutral cell cache: each sequence's logits must equal running it on + # its own -- per-token RoPE, windows and caches all kept apart. + generator = torch.Generator().manual_seed(1) + a = torch.randint(0, TINY_CONFIG.vocab_size, (1, 20), generator=generator) + b = torch.randint(0, TINY_CONFIG.vocab_size, (1, 7), generator=generator) + a_next = torch.randint(0, TINY_CONFIG.vocab_size, (1, 1), generator=generator) + b_next = torch.randint(0, TINY_CONFIG.vocab_size, (1, 1), generator=generator) + + def alone(prompt, next_token): + key = self._install( + SequenceReferenceCache( + _cache_config(self.off_graph, TINY_CONFIG.max_seq_len) + ) + ) + length = prompt.shape[1] + with torch.no_grad(), REGISTRY.active(key): + first = step_forward( + self.off_graph, + prompt, + torch.arange(length), + torch.tensor([length - 1]), + ) + second = step_forward( + self.off_graph, + next_token, + torch.tensor([length]), + torch.tensor([0]), + ) + return first[0], second[0] + + a_first, a_second = alone(a, a_next) + b_first, b_second = alone(b, b_next) + + cells = CellReferenceCache(_cache_config(self.off_graph, 64)) + key = self._install(cells) + with torch.no_grad(), REGISTRY.active(key): + cells.declare_step([0] * 20 + [1] * 7) + prefill = step_forward( + self.off_graph, + torch.cat([a, b], dim=1), + torch.cat([torch.arange(20), torch.arange(7)]), + torch.tensor([19, 26]), + ) + cells.declare_step([1, 0]) + decode = step_forward( + self.off_graph, + torch.cat([b_next, a_next], dim=1), + torch.tensor([7, 20]), + torch.tensor([0, 1]), + ) + for actual, expected in ( + (prefill[0], a_first), + (prefill[1], b_first), + (decode[1], a_second), + (decode[0], b_second), + ): + self.assertLess((expected - actual).abs().max().item(), 5e-2) + self.assertEqual(int(expected.argmax()), int(actual.argmax())) + + def test_cell_manifest_keeps_every_layer(self) -> None: + model = build_random_tiny_model() + manifest = json.loads(cell_manifest(enable_offgraph_kv_cache(model, 8), 128)) + self.assertEqual(manifest["layout"], "cell") + self.assertEqual(manifest["max_cells"], 128) + self.assertEqual(manifest["max_write"], 8) + self.assertEqual(len(manifest["layers"]), TINY_CONFIG.n_layers) + + +class BatchingExportTest(unittest.TestCase): + MAX_STEP = 16 + MAX_CELLS = 128 + + def setUp(self) -> None: + if not torch.cuda.is_available(): + self.skipTest("CUDA required") + + def test_exports_decode_and_prefill_over_a_cell_cache(self) -> None: + from executorch.backends.cuda.passes.lower_offgraph_kv import ( + OFFGRAPH_KV_COMPILE_SPEC, + OFFGRAPH_KV_STEP_WIDTH_COMPILE_SPEC, + ) + from executorch.exir.lowered_backend_module import get_lowered_submodules + + with ( + tempfile.TemporaryDirectory() as ckpt_dir, + tempfile.TemporaryDirectory() as out_dir, + ): + save_checkpoint(ckpt_dir) + model, config = load_prequantized_model( + ckpt_dir, max_seq_len=TINY_CONFIG.max_seq_len + ) + program = export_batching( + model, + config, + out_dir, + max_step_tokens=self.MAX_STEP, + max_cells=self.MAX_CELLS, + ) + self.assertTrue(os.path.exists(os.path.join(out_dir, "model.pte"))) + self.assertTrue( + any(name.endswith(".ptd") for name in os.listdir(out_dir)) + ) + + methods = set(program.methods) + self.assertTrue({"decode", "prefill"}.issubset(methods)) + self.assertFalse( + {"embed_text", "forward_from_embeddings", "decode_from_embedding"} + & methods + ) + for name in ("decode", "prefill"): + graph = program.exported_program(name).graph_module + lowered = get_lowered_submodules(graph) + self.assertEqual(len(lowered), 1, name) + specs = { + spec.key: spec.value for spec in lowered[0][1].compile_specs + } + manifest = json.loads(specs[OFFGRAPH_KV_COMPILE_SPEC]) + self.assertEqual(manifest["layout"], "cell") + self.assertEqual(manifest["max_cells"], self.MAX_CELLS) + self.assertEqual(manifest["max_write"], self.MAX_STEP) + # The step width must name the delegate input carrying T, which + # the backend checks is the cache ops' position. The delegate's + # inputs need not follow the method's order: prefill passes + # logits_to_keep ahead of input_pos. + index, dim = map( + int, specs[OFFGRAPH_KV_STEP_WIDTH_COMPILE_SPEC].decode().split(":") + ) + call = next( + n + for n in graph.graph.nodes + if "executorch_call_delegate" in str(n.target) + ) + tokens = next( + n for n in graph.graph.nodes if n.op == "placeholder" + and n.name == "tokens" + ) + self.assertEqual( + str(call.args[1 + index].meta["val"].shape[dim]), + str(tokens.meta["val"].shape[1]), + name, + ) + + from executorch.runtime import Runtime, Verification + + loaded = Runtime.get().load_program( + os.path.join(out_dir, "model.pte"), verification=Verification.Minimal + ) + + def constant(name): + return loaded.load_method(name).execute([])[0] + + self.assertEqual(constant(MAX_CELLS_METHOD), self.MAX_CELLS) + self.assertEqual(constant("get_max_context_len"), TINY_CONFIG.max_seq_len) + self.assertEqual(constant("get_n_caches"), TINY_CONFIG.n_layers) + self.assertEqual(constant("get_max_prefill_chunk"), self.MAX_STEP) + self.assertEqual(constant("get_min_prefill_chunk"), 5) + # LogitsToKeepMode::Selected. + self.assertEqual(constant("get_logits_to_keep_mode"), 2)