From 8fe962fa207f002c6cc6abcece7cf4a9ca55fae2 Mon Sep 17 00:00:00 2001 From: Joshua Calafato Date: Wed, 16 Sep 2026 23:29:31 +0000 Subject: [PATCH 1/6] feat(gemma): add paired ONNX execution Forward explicit Gemma4 assistant MTP and DSpark block7 pairs to the pinned native Edge-LLM exporter, builder and runtime. Keep pair admission, artifacts and generation controls family-owned. Render the checkpoint single-user template before Edge tokenization to preserve disabled-thinking behavior without output filtering. Retain meaningful quality gates and document the exact locally qualified text profiles and remaining CI registration gap. Signed-off-by: Joshua Calafato --- families/gemma/EDGE_LLM.md | 95 +++++++ families/gemma/edge_llm.py | 241 ++++++++++++++++ families/gemma/model.py | 7 + families/gemma/runtime/CMakeLists.txt | 22 +- families/gemma/runtime/edge_llm/adapter.cpp | 257 ++++++++++++++++++ families/gemma/runtime/edge_llm/adapter.h | 15 + families/gemma/runtime/edge_llm/contract.h | 57 ++++ .../gemma/runtime/edge_llm/device_link.cu | 5 + families/gemma/runtime/edge_llm/request.h | 68 +++++ families/gemma/runtime/plugin.cpp | 10 + families/gemma/support.py | 2 +- .../gemma/tests/cpp/test_gemma_sampler.cpp | 36 +++ families/gemma/tests/test_e2e.py | 5 +- families/gemma/tests/test_model_type_gate.py | 6 +- website/docs/features/model-families.md | 17 ++ 15 files changed, 836 insertions(+), 7 deletions(-) create mode 100644 families/gemma/EDGE_LLM.md create mode 100644 families/gemma/edge_llm.py create mode 100644 families/gemma/runtime/edge_llm/adapter.cpp create mode 100644 families/gemma/runtime/edge_llm/adapter.h create mode 100644 families/gemma/runtime/edge_llm/contract.h create mode 100644 families/gemma/runtime/edge_llm/device_link.cu create mode 100644 families/gemma/runtime/edge_llm/request.h diff --git a/families/gemma/EDGE_LLM.md b/families/gemma/EDGE_LLM.md new file mode 100644 index 0000000000..91cc8dbf53 --- /dev/null +++ b/families/gemma/EDGE_LLM.md @@ -0,0 +1,95 @@ +# Gemma4 MTP Edge-LLM adapter + +This family owns the explicit Gemma4 unified target/assistant pair, ONNX +command mapping, bundle assets and C++ runtime adapter. Standalone Gemma1/2 +builds remain native. Selecting a Gemma4 checkpoint alone does not enable this +paired path or claim native Gemma4 support. + +## Build and inference + +Use the [pinned native SDK provisioning](../../cmake/edgellm/README.md) with +`TRTMC_EDGELLM_ALL_KERNELS=ON` and `TRTMC_EDGELLM_ONNX=ON`, then configure the +Model Connect runtime with `TRTMC_ENABLE_EDGELLM=ON`. Set `CMAKE_PREFIX_PATH` +to the SDK installation. The source is official GitHub Edge-LLM 0.10.1, +revision `e8b29522938901f6df19ebeedd4b69bc8edbcd97`; cross compilation is not used. + +Pass `execution=BuildExecutionInputs("mtp", (NamedCheckpoint("draft", +draft_path),))` to the Python build API. The build CLI equivalent adds +`--execution-variant mtp --companion draft=/path/to/assistant`. +The family forwards both unmodified checkpoints to the original Python ONNX +exporter with `--mtp --mtp-draft-dir`, excluding image/audio branches for this +text-only profile. The original native ONNX builder creates both speculative +engines. Those engines, embedding and tokenizer assets are bundled; source +checkpoint weights and ONNX intermediates are not duplicated in the bundle. + +The C++ adapter owns a persistent original Edge inference runtime with drafting +topK 1, steps 3 and verify size 4. This upstream MTP runtime uses greedy decoding; +the adapter rejects a requested non-greedy configuration instead of silently +changing it. Unmapped generation controls and capacity overflow are rejected. + +An Edge preparation failure emits a warning and retains diagnostics. Native +Gemma currently cannot implement this MTP request, so the build fails explicitly +rather than substituting an ordinary base-only engine. Runtime errors propagate +without fallback. + +## Qualification scope + +The completed MTP qualification uses: + +- Target `google/gemma-4-12B-it`, revision + `707f0a3b8a3c7ad586ed01e27eafbad8a27dd0f7`. +- Assistant `google/gemma-4-12B-it-assistant`, revision + `46d4c6f13f0ac0ad827b915669b8df9b81c64c51`. +- FP16 text execution, native SM80, CUDA 13.3, TensorRT 11.1.0.106, + TP1/batch1, input 512 and KV capacity 1024, greedy chat with thinking disabled. +- A fresh independent BF16 Hugging Face reference and unchanged Model Connect + normalized edit distance gate of 0.15. +- Original Edge `llm_basic` prompt and 128-token budget, with unchanged ROUGE-1 + / ROUGE-L gates of 0.25 / 0.20. Explicit greedy controls match upstream MTP's + effective behavior; they do not claim parity with a sampling request. + +The actual MTP Model Connect build and public CLI inference pass. Fresh HF +reference tokens match exactly (NED **0.0** against **0.15**). The 128-token +fixture passes ROUGE-1 **0.4343** / ROUGE-L **0.2286** against **0.25 / 0.20**. +Both checks were repeated successfully on this combined MTP/DSpark runtime, +using the same MTP bundle and a source-verified reused reference. Native +compilation, both C++ tests and all 10 existing family Python checks pass. +This recipe does not qualify other models, platforms, multimodal inputs or +sampling. It reuses existing family E2E helpers; the pair is not yet a registered +pytest manifest case. + +## Qualified DSpark block7 extension + +The same owning family also admits `--execution-variant dspark` with target +`google/gemma-4-12B-it` at the revision above and draft +`deepseek-ai/dspark_gemma4_12b_block7` at +`2fa72e765eec2965fc4d86a8663ce6769eba6218`. It forwards original +`--dspark-base` and `--dspark-draft` exports, packages both speculative engines +and confidence/Markov sidecars, and uses drafting topK 1, step 1, verify 8, +block7, scheduler off. Supported sampling controls remain enabled for DSpark; +the MTP-only greedy restriction is not applied to this variant. + +Actual DSpark Model Connect ONNX export/build and public CLI inference pass on +SM80 with the same capacities above. Fresh independent BF16 HF reference tokens +match exactly (NED **0.0** <= **0.15**). The original 128-token Edge fixture uses +temperature 1, topK 50 and topP 1: ROUGE-1 **0.4114** / ROUGE-L **0.2286** pass +unchanged **0.25 / 0.20** gates. Both existing C++ tests and all 10 existing +family Python tests pass. This proves this exact sampled run, not statistical +sampling equivalence or other model/platform combinations. + + +## Checkpoint chat-template correction + +The official Edge 0.10.1 static Gemma template omits the checkpoint's closed +thought channel when thinking is disabled, and differs in enabled-thinking +system-prefix handling. A first uncorrected MTP run produced `thought\nParis` +instead of the independent reference `Paris`, failing NED 0.6154 against 0.15. +The family now validates the source single-user template during build and +renders it faithfully before calling Edge with raw text. Unicode whitespace +trimming follows the checkpoint Jinja filter; raw-text requests are unchanged. +No generated-output filtering, engine change or relaxed quality gate is used. +The original Edge fixture is scored with this source-faithful prompt mapping, +not claimed as byte-for-byte parity with the faulty upstream static template. +MTP passes both unchanged quality gates with the same engine bundle after the +request-only fix, including on the combined runtime. DSpark also passes both +unchanged quality gates with the source-faithful prompt mapping. Initial failure evidence and the tokenizer audit are retained. diff --git a/families/gemma/edge_llm.py b/families/gemma/edge_llm.py new file mode 100644 index 0000000000..8eb880d9e8 --- /dev/null +++ b/families/gemma/edge_llm.py @@ -0,0 +1,241 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Family-owned Gemma4 paired forwarding to the pinned original ONNX toolchain.""" + +from __future__ import annotations + +import json +import logging +import os +from pathlib import Path +import shutil +import subprocess +import tempfile +import traceback + +from tensorrt_model_connect.build import cmake_prefixes, detect_local_platform, subprocess_environment + +EDGE_REVISION = "e8b29522938901f6df19ebeedd4b69bc8edbcd97" + + +def local_target() -> dict: + """Return the executing worker identity supplied by generic build mechanics.""" + return detect_local_platform() + + +def installed_package(target: dict) -> dict: + """Resolve CMake installation via standard prefixes; never install anything. + + Args: + target: Executing device and SDK identity. + + Returns: + Validated package metadata with absolute Python and plugin paths. + + Raises: + FileNotFoundError: No CMake installation or required artifact exists. + ValueError: Pin, architecture, SDK or contained-path contract differs. + """ + for prefix in cmake_prefixes(): + manifest = prefix / "share/trtmc/edge-llm.json" + if not manifest.is_file(): + continue + package = json.loads(manifest.read_text(encoding="utf-8")) + if package.get("schema_version") != 1 or package.get("revision") != EDGE_REVISION: + raise ValueError(f"Edge package has an unsupported revision/schema: {manifest}") + if package.get("version") != "0.10.1" or package.get("arch") != target["arch"]: + raise ValueError("Edge package version/architecture differs from executing worker") + if target["sm"] not in package.get("architectures", []): + raise ValueError("Edge package was not built for this local GPU") + cuda_version = ".".join(str(package.get("cuda_version", "")).split(".")[:2]) + if cuda_version != target["cuda_version"] or package.get("tensorrt_version") != target["tensorrt_version"]: + raise ValueError("Edge package CUDA/TensorRT differs from executing worker") + for name in ("python", "plugin") + (("onnx_builder",) if package.get("onnx") else ()): + relative = Path(package[name]) + path = (prefix / relative).resolve() + if relative.is_absolute() or not path.is_relative_to(prefix.resolve()): + raise ValueError(f"Edge package {name} must be contained in its installation") + if not path.is_file(): + raise FileNotFoundError(f"Edge package {name} is missing: {path}") + package[name] = str(path) + return package + raise FileNotFoundError("Edge-LLM is not installed; enable the optional Edge-LLM CMake dependency " + "and set CMAKE_PREFIX_PATH to its install prefix") + + +_LOG = logging.getLogger(__name__) + +# Validate the exact source template; the runtime owns its single-user rendering. +_CHAT_ADMISSION = r""" +from transformers import AutoTokenizer +import sys +tokenizer = AutoTokenizer.from_pretrained(sys.argv[1], local_files_only=True, trust_remote_code=False) +for thinking in (False, True): + for text in ("hello", " hello\n", "\u2003hello\u3000"): + expected = "" + if thinking: + expected += "<|turn>system\n<|think|>\n\n" + expected += "<|turn>user\n" + text.strip() + "\n<|turn>model\n" + if not thinking: + expected += "<|channel>thought\n" + actual = tokenizer.apply_chat_template([{"role": "user", "content": text}], + tokenize=False, add_generation_prompt=True, enable_thinking=thinking) + if actual != expected: + raise ValueError("Gemma4 source chat template differs from the family runtime mapping") +print("Gemma4 single-user chat template admission passed") +""" + + +def validate_pair(request, execution) -> tuple[dict, Path, int]: + """Admit only explicit unquantized Gemma4 unified assistant text execution.""" + if execution.variant not in {"mtp", "dspark"} or tuple(c.role for c in execution.checkpoints) != ("draft",): + raise ValueError("Gemma4 pairs require variant=mtp or dspark and one named draft checkpoint") + if (request.backend != "trt" or request.task != "text_generation" + or request.precision != "fp16" or request.quantization not in {None, "none"} + or request.max_batch_size != 1 or request.tensor_parallel_size != 1 + or request.context_parallel_size != 1 or request.dynamic_kv_cache + or request.fp32_layers or request.graph_transform is not None + or any(v is not None for v in (request.image_height, request.image_width, + request.video_num_frames))): + raise ValueError("Gemma4 paired execution maps FP16 text-only TP1/batch1 without graph transforms") + source, draft = Path(request.model_dir), execution.checkpoints[0].model_dir + raw = json.loads((source / "config.json").read_text()) + companion = json.loads((draft / "config.json").read_text()) + if not isinstance(raw, dict) or raw.get("model_type") != "gemma4_unified": + raise ValueError("Gemma4 MTP expects a Gemma4 unified target") + if not isinstance(companion, dict): + raise ValueError("Gemma4 requires a companion configuration object") + base = raw.get("text_config") + if not isinstance(base, dict): + raise ValueError("Gemma4 target requires nested text configuration") + if execution.variant == "mtp": + if companion.get("model_type") != "gemma4_unified_assistant": + raise ValueError("Gemma4 MTP expects a unified assistant") + assistant = companion.get("text_config") + if not isinstance(assistant, dict) or companion.get("backbone_hidden_size") != base.get("hidden_size"): + raise ValueError("Gemma4 MTP target and assistant geometry is incompatible") + else: + assistant = companion + if (companion.get("architectures") != ["Gemma4DSparkModel"] + or companion.get("target_model_type") != "gemma4_unified" + or companion.get("hidden_size") != base.get("hidden_size") + or companion.get("num_target_layers") != base.get("num_hidden_layers") + or companion.get("block_size") != 7): + raise ValueError("Gemma4 DSpark requires a matching target and block7 draft") + layers = companion.get("target_layer_ids") + if (not isinstance(layers, list) or not layers + or any(type(i) is not int or not 0 <= i < base["num_hidden_layers"] for i in layers) + or len(set(layers)) != len(layers)): + raise ValueError("Invalid Gemma4 DSpark target layer IDs") + mask = companion.get("mask_token_id") + if type(mask) is not int or not 0 <= mask < base["vocab_size"]: + raise ValueError("Invalid Gemma4 DSpark mask token") + if (assistant.get("vocab_size") != base.get("vocab_size") + or base.get("enable_moe_block") or assistant.get("enable_moe_block")): + raise ValueError("Gemma4 pair vocabulary or dense topology is incompatible") + for directory, config in ((source, raw), (draft, companion)): + if (config.get("quantization_config") or config.get("text_config", config).get("quantization_config") + or any((directory / name).exists() for name in + ("hf_quant_config.json", "quantize_config.json", "quant_config.json"))): + raise ValueError("This Gemma4 paired profile requires unquantized checkpoints") + if not list(directory.glob("*.safetensors")): + raise ValueError("Gemma4 paired execution requires both local safetensors checkpoints") + limit = request.max_sequence_length or 1024 + capacities = (base.get("max_position_embeddings"), assistant.get("max_position_embeddings")) + if any(type(v) is not int or v < limit for v in capacities) or not (4 if execution.variant == "mtp" else 8) < limit <= 1024: + raise ValueError("Gemma4 pair requires context above its verify size and at most 1024") + return raw, draft, limit + + +def prepare(request, raw, draft: Path, limit: int, target: dict, staging: Path, log_path: Path, variant: str): + """Forward exact checkpoints to the original exporter and native builder.""" + package = installed_package(target) + if package.get("onnx") is not True: + raise ValueError("Gemma4 paired execution requires an ONNX-enabled Edge SDK") + source = Path(request.model_dir).resolve() + engine, onnx = staging / "edge_llm/engine", staging / "onnx" + checkpoint = staging / "edge_llm/checkpoint" + checkpoint.mkdir(parents=True) + shutil.copy2(source / "config.json", checkpoint / "config.json") + (checkpoint / "draft").mkdir() + shutil.copy2(draft / "config.json", checkpoint / "draft/config.json") + env = subprocess_environment( + {"EDGELLM_PLUGIN_PATH": package["plugin"]}, + prepend_paths={"LD_LIBRARY_PATH": str(Path(package["plugin"]).parent)}, + ) + exporter = [package["python"], "-I", "-m", "tensorrt_edgellm.scripts.export", + str(source), str(onnx)] + if variant == "mtp": + commands = [exporter + ["--mtp", "--mtp-draft-dir", str(draft.resolve()), + "--skip-visual", "--skip-audio"]] + draft_subdir, verify_size, draft_size, spec_type = "mtp_draft", 4, 4, "gemma4_mtp" + else: + commands = [exporter + [flag, "--dspark-draft-dir", str(draft.resolve()), + "--skip-visual", "--skip-audio"] + for flag in ("--dspark-base", "--dspark-draft")] + draft_subdir, verify_size, draft_size, spec_type = "dspark_draft", 8, 7, "dspark" + for subdirectory, flag in (("llm", "--specBase"), (draft_subdir, "--specDraft")): + commands.append([package["onnx_builder"], "--onnxDir", str(onnx / subdirectory), + "--engineDir", str(engine), flag, "--maxInputLen", str(min(limit, 512)), + "--maxKVCacheCapacity", str(limit), "--maxBatchSize", "1", + "--maxVerifyTreeSize", str(verify_size), "--maxDraftTreeSize", str(draft_size)]) + commands.insert(0, [package["python"], "-I", "-c", _CHAT_ADMISSION, str(source)]) + with log_path.open("a", encoding="utf-8") as log: + for command in commands: + log.write(json.dumps(command) + "\n") + log.flush() + subprocess.run(command, check=True, stdout=log, stderr=subprocess.STDOUT, + cwd=staging, env=env) + required = ("spec_base.engine", "spec_draft.engine", "base_config.json", "draft_config.json", + "embedding.safetensors", "tokenizer.json", "tokenizer_config.json", + "processed_chat_template.json") + if variant == "dspark": + required += ("dspark_heads.safetensors", "dspark_heads_info.json") + for name in required: + if not (engine / name).is_file() or (engine / name).stat().st_size == 0: + raise ValueError(f"Edge Gemma4 MTP builder omitted required artifact: {name}") + for role in ("base", "draft"): + config = json.loads((engine / f"{role}_config.json").read_text()) + if config.get("spec_decode_type") != spec_type: + raise ValueError("Edge returned a different Gemma4 speculative contract") + files = {} + for directory in (engine, checkpoint): + for path in sorted(directory.rglob("*")): + if path.is_symlink(): + raise ValueError(f"Gemma4 Edge output must not contain symlinks: {path}") + if path.is_file(): + files[path.relative_to(staging).as_posix()] = path + return files, { + "version": 1, "edge_revision": EDGE_REVISION, "target": target, "precision": "fp16", + "max_sequence_length": limit, "max_input_length": min(limit, 512), "max_batch_size": 1, + "execution_variant": variant, "builder_flow": "onnx", "artifacts": list(files), + } + + +def build(request, writer, execution) -> None: + """Prepare the full pair before publication; never fall back to a base-only engine.""" + raw, draft, limit = validate_pair(request, execution) + descriptor, name = tempfile.mkstemp(prefix=f".{request.output_path.name}.edge-onnx-", + suffix=".log", dir=request.output_path.parent) + os.close(descriptor) + log_path = Path(name) + with tempfile.TemporaryDirectory(prefix="trtmc-gemma4-edge-") as directory: + try: + target = local_target() + if (target["os"], target["arch"], target["sm"]) != ("linux", "x86_64", 80): + raise ValueError("This Gemma4 MTP Edge profile maps native x86_64 SM80") + files, marker = prepare(request, raw, draft, limit, target, Path(directory), log_path, execution.variant) + except Exception as error: + with log_path.open("a", encoding="utf-8") as log: + traceback.print_exception(error, file=log) + _LOG.warning("Gemma4 Edge build failed: %s. Diagnostics: %s. " + "Native fallback cannot preserve the requested paired variant.", error, log_path) + # The native family has no paired Gemma4 topology; do not substitute + # an autoregressive decoder and misrepresent execution semantics. + raise NotImplementedError("Native Gemma does not implement the requested paired variant") from error + writer.set_header(family="gemma", task=request.task, backend=request.backend) + for name, path in files.items(): + with path.open("rb") as source, writer.open_section(name) as destination: + shutil.copyfileobj(source, destination, length=1024 * 1024) + writer.add_json("edge_llm.json", marker) diff --git a/families/gemma/model.py b/families/gemma/model.py index 2ab719e1c2..8418906039 100644 --- a/families/gemma/model.py +++ b/families/gemma/model.py @@ -532,3 +532,10 @@ def build(request: "BuildRequest", writer: "BundleWriter") -> None: path = model_dir / filename if path.is_file(): writer.add_bytes(filename, path.read_bytes()) + + +def build_with_inputs(request, writer, execution) -> None: + """Forward the explicit Gemma4 MTP or DSpark pair without changing standalone builds.""" + from .edge_llm import build as build_pair + + build_pair(request, writer, execution) diff --git a/families/gemma/runtime/CMakeLists.txt b/families/gemma/runtime/CMakeLists.txt index 5303181cd9..cc138d7716 100644 --- a/families/gemma/runtime/CMakeLists.txt +++ b/families/gemma/runtime/CMakeLists.txt @@ -26,7 +26,10 @@ target_link_libraries(trtmc_model_gemma PRIVATE ${TRTMC_CUDART_LIBRARY} ${CMAKE_DL_LIBS} ) -target_compile_options(trtmc_model_gemma PRIVATE -Wall -Wextra -Wpedantic) +target_compile_options(trtmc_model_gemma PRIVATE + "$<$:-Wall;-Wextra;-Wpedantic>" + "$<$:-Xcompiler=-Wall,-Wextra>" +) set_target_properties(trtmc_model_gemma PROPERTIES LIBRARY_OUTPUT_DIRECTORY "${CMAKE_BINARY_DIR}" BUILD_RPATH "\$ORIGIN" @@ -54,3 +57,20 @@ if(TRTMC_BUILD_TESTS) endforeach() set_tests_properties(test_gemma_pipeline PROPERTIES SKIP_RETURN_CODE 77) endif() + +# Optional complete-network offload; all model-specific orchestration stays here. +if(TARGET EdgeLLM::Core) + target_sources(trtmc_model_gemma PRIVATE edge_llm/adapter.cpp edge_llm/device_link.cu) + target_compile_definitions(trtmc_model_gemma PRIVATE TRTMC_HAS_EDGE_LLM=1) + target_link_libraries(trtmc_model_gemma PRIVATE EdgeLLM::Core) + set_target_properties(trtmc_model_gemma PROPERTIES + CUDA_ARCHITECTURES "${EdgeLLM_CUDA_ARCHITECTURE}" + CUDA_SEPARABLE_COMPILATION ON CUDA_RESOLVE_DEVICE_SYMBOLS ON) + add_custom_command(TARGET trtmc_model_gemma POST_BUILD + COMMAND ${CMAKE_COMMAND} -E copy_if_different + $ $ VERBATIM) + if(TRTMC_BUILD_TESTS) + target_compile_definitions(test_gemma_sampler PRIVATE TRTMC_HAS_EDGE_LLM=1) + target_link_libraries(test_gemma_sampler PRIVATE EdgeLLM::Core) + endif() +endif() diff --git a/families/gemma/runtime/edge_llm/adapter.cpp b/families/gemma/runtime/edge_llm/adapter.cpp new file mode 100644 index 0000000000..45cb0626c6 --- /dev/null +++ b/families/gemma/runtime/edge_llm/adapter.cpp @@ -0,0 +1,257 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#include "families/gemma/runtime/edge_llm/adapter.h" + +#include "families/gemma/runtime/edge_llm/request.h" + +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +namespace trtmc::gemma::edge_llm { +namespace { +namespace fs = std::filesystem; + +/// Turn CUDA failures into caller-visible load or inference errors. +void check_cuda(cudaError_t result) { + if (result != cudaSuccess) + throw std::runtime_error(std::string("Gemma4 Edge CUDA error: ") + + cudaGetErrorString(result)); +} + +/// Reject an engine built for a different local GPU or CUDA/TensorRT runtime. +void validate_target(const nlohmann::json& target) { + utsname host{}; + if (uname(&host) != 0) + throw std::runtime_error("Cannot identify Gemma4 Edge runtime host"); + std::ifstream release("/etc/os-release"); + std::string line, os_version; + while (std::getline(release, line)) { + if (line.rfind("VERSION_ID=", 0) == 0) { + os_version = line.substr(11); + if (os_version.size() >= 2 && os_version.front() == char(34) && + os_version.back() == char(34)) + os_version = os_version.substr(1, os_version.size() - 2); + } + } + int device = 0, cuda_version = 0; + check_cuda(cudaGetDevice(&device)); + check_cuda(cudaRuntimeGetVersion(&cuda_version)); + cudaDeviceProp gpu{}; + check_cuda(cudaGetDeviceProperties(&gpu, device)); + const int trt_version = getInferLibVersion(); + const std::string trt = + std::to_string(trt_version / 10000) + "." + std::to_string((trt_version % 10000) / 100) + + "." + std::to_string(trt_version % 100) + "." + std::to_string(getInferLibBuildVersion()); + const std::string cuda = + std::to_string(cuda_version / 1000) + "." + std::to_string((cuda_version % 1000) / 10); + if (target.at("os") != "linux" || target.at("os_version") != os_version || + target.at("arch") != host.machine || target.at("sm") != gpu.major * 10 + gpu.minor || + target.at("cuda_version") != cuda || target.at("tensorrt_version") != trt) + throw std::runtime_error( + "Gemma4 Edge bundle requires its build GPU and CUDA/TensorRT stack"); +} + +/// Own extracted engine/checkpoint files until after the Edge runtime is destroyed. +class Artifacts { + public: + explicit Artifacts(const BundleReader& bundle, const nlohmann::json& marker) { + std::set names; + for (const auto& entry : marker.at("artifacts")) { + const auto name = entry.get(); + if (!safe_artifact_path(name) || !names.insert(name).second || + !bundle.find_section(name)) + throw std::runtime_error("Invalid Gemma4 Edge artifact: " + name); + } + std::vector required_files{"edge_llm/engine/tokenizer.json", + "edge_llm/engine/tokenizer_config.json", + "edge_llm/engine/processed_chat_template.json", + "edge_llm/engine/embedding.safetensors", + "edge_llm/engine/spec_base.engine", + "edge_llm/engine/spec_draft.engine", + "edge_llm/engine/base_config.json", + "edge_llm/engine/draft_config.json"}; + if (marker.at("execution_variant") == "dspark") { + required_files.push_back("edge_llm/engine/dspark_heads.safetensors"); + required_files.push_back("edge_llm/engine/dspark_heads_info.json"); + } + for (const auto& required : required_files) + if (!names.count(required) || bundle.find_section(required)->length == 0) + throw std::runtime_error("Required Gemma4 Edge artifact missing: " + required); + std::string pattern = (fs::temp_directory_path() / "trtmc-gemma-edge-XXXXXX").string(); + if (!mkdtemp(pattern.data())) + throw std::runtime_error("Cannot create Gemma4 Edge artifact directory"); + root_ = pattern; + try { + for (const auto& name : names) { + const auto destination = root_ / name; + fs::create_directories(destination.parent_path()); + std::ofstream output(destination, std::ios::binary); + bundle.copy_section(name, output); + output.close(); + if (!output) + throw std::runtime_error("Cannot extract Gemma4 Edge artifact: " + name); + } + } catch (...) { + cleanup(); + throw; + } + } + ~Artifacts() { cleanup(); } + Artifacts(const Artifacts&) = delete; + Artifacts& operator=(const Artifacts&) = delete; + std::string engine() const { return (root_ / "edge_llm/engine").string(); } + + private: + void cleanup() noexcept { + std::error_code ignored; + fs::remove_all(root_, ignored); + } + fs::path root_; +}; + +/// Close the plugin handle after runtime destruction; registrations remain mapped. +struct CloseLibrary { + void operator()(void* handle) const noexcept { + if (handle) + dlclose(handle); + } +}; + +/// Initialize the CMake-installed adjacent plugin without process-global environment mutation. +std::unique_ptr load_plugin() { + Dl_info location{}; + if (!dladdr(reinterpret_cast(&create), &location) || !location.dli_fname) + throw std::runtime_error("Cannot locate Gemma4 family library"); + const auto path = + fs::absolute(location.dli_fname).parent_path() / "libNvInfer_edgellm_plugin.so"; + std::unique_ptr plugin( + dlopen(path.c_str(), RTLD_NOW | RTLD_GLOBAL | RTLD_NODELETE)); + if (!plugin) + throw std::runtime_error("Cannot load CMake-installed Edge plugin: " + + std::string(dlerror())); + using Initialize = bool (*)(void*, const char*); + auto initialize = reinterpret_cast(dlsym(plugin.get(), "initEdgellmPlugins")); + if (!initialize || !initialize(static_cast(&trt_edgellm::gLogger), "")) + throw std::runtime_error("Cannot initialize Gemma4 Edge plugin"); + return plugin; +} + +/// Stream ownership is independent of construction success and outlives the Edge instance. +class Stream { + public: + Stream() { check_cuda(cudaStreamCreateWithFlags(&value_, cudaStreamNonBlocking)); } + ~Stream() { cudaStreamDestroy(value_); } + Stream(const Stream&) = delete; + Stream& operator=(const Stream&) = delete; + cudaStream_t get() const { return value_; } + + private: + cudaStream_t value_{nullptr}; +}; + +/// Delegate the complete assistant MTP algorithm to the original Edge runtime. +std::unique_ptr +make_runtime(const Artifacts& artifacts, cudaStream_t stream, bool dspark) { + trt_edgellm::rt::SpecDecodeDraftingConfig drafting{}; + drafting.draftingTopK = 1; + drafting.draftingStep = dspark ? 1 : 3; + drafting.verifySize = dspark ? 8 : 4; + drafting.dsparkSchedulerMode = trt_edgellm::rt::DSparkSchedulerMode::kOff; + return std::make_unique( + artifacts.engine(), "", std::unordered_map{}, drafting, stream, + trt_edgellm::rt::ContextCacheConfig{}); +} + +/// Thin persistent Edge API adapter; serialization prevents concurrent use of Edge request state. +class EdgeTask final : public ITextGeneration { + public: + EdgeTask(const BundleReader& bundle, const nlohmann::json& marker) + : artifacts_(bundle, marker), plugin_(load_plugin()), + runtime_( + make_runtime(artifacts_, stream_.get(), marker.at("execution_variant") == "dspark")), + capacity_(marker.at("max_sequence_length").get()), + input_limit_(marker.at("max_input_length").get()), + dspark_(marker.at("execution_variant") == "dspark") {} + + std::int32_t default_max_new_tokens() const override { return std::min(128, capacity_ - 1); } + + /// Drain work from failed requests before destroying the runtime and its weight buffers. + ~EdgeTask() override { cudaStreamSynchronize(stream_.get()); } + + /// Invoke Edge once; failures propagate without attempting native inference. + TextResult generate(const std::string& prompt, const TextGenerationConfig& config) override { + auto request = make_request(prompt, config, default_max_new_tokens(), dspark_); + std::lock_guard lock(mutex_); + const auto counts = runtime_->countPromptTokens(request); + if (counts.size() != 1) + throw std::runtime_error("Gemma4 Edge returned invalid prompt counts"); + validate_capacity(counts.front(), input_limit_, capacity_, request.maxGenerateLength); + trt_edgellm::rt::LLMGenerationResponse response{}; + // Complete queued work before response/request storage is destroyed, + // including exception paths in a persistent task. + struct Drain { + cudaStream_t stream; + ~Drain() { cudaStreamSynchronize(stream); } + } drain{stream_.get()}; + if (!runtime_->handleRequest(request, response, stream_.get()) || + response.outputIds.size() != 1 || response.outputTexts.size() != 1 || + response.outputIds.front().empty() || + response.outputIds.front().size() > static_cast(request.maxGenerateLength)) + throw std::runtime_error("Gemma4 Edge generation failed"); + if (response.finishReasons.size() != 1 || + (response.finishReasons.front() != trt_edgellm::rt::FinishReason::kEndId && + response.finishReasons.front() != trt_edgellm::rt::FinishReason::kLength)) + throw std::runtime_error("Gemma4 Edge generation did not complete successfully"); + // This API does not expose per-request stage times; zero means unavailable. + return {std::move(response.outputTexts.front()), std::move(response.outputIds.front())}; + } + + private: + // Reverse destruction order keeps weights, plugin and stream alive throughout Edge teardown. + Artifacts artifacts_; + std::unique_ptr plugin_; + Stream stream_; + std::unique_ptr runtime_; + int capacity_; + int input_limit_; + bool dspark_; + std::mutex mutex_; +}; +} // namespace + +ITask* create(const BundleReader& bundle) { + const auto bytes = bundle.read_section("edge_llm.json"); + const auto marker = nlohmann::json::parse(bytes.begin(), bytes.end()); + if (marker.at("version") != 1 || marker.at("edge_revision") != kRevision || + marker.at("max_sequence_length").get() <= 1 || + marker.at("max_input_length").get() <= 0 || + marker.at("max_input_length").get() > marker.at("max_sequence_length").get() || + marker.at("max_batch_size") != 1 || marker.at("precision") != "fp16" || + !marker.at("artifacts").is_array()) + throw std::runtime_error("Invalid Gemma4 Edge bundle contract"); + if ((marker.value("execution_variant", "") != "mtp" && + marker.value("execution_variant", "") != "dspark") || + marker.value("builder_flow", "") != "onnx") + throw std::runtime_error("Gemma4 requires a paired MTP or DSpark ONNX contract"); + validate_target(marker.at("target")); + return new EdgeTask(bundle, marker); +} + +} // namespace trtmc::gemma::edge_llm diff --git a/families/gemma/runtime/edge_llm/adapter.h b/families/gemma/runtime/edge_llm/adapter.h new file mode 100644 index 0000000000..5f4d6c53e9 --- /dev/null +++ b/families/gemma/runtime/edge_llm/adapter.h @@ -0,0 +1,15 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#pragma once + +#include "trtmc/bundle.h" +#include "trtmc/task.h" + +namespace trtmc::gemma::edge_llm { + +/// Create a persistent Edge task from a self-contained bundle; throws on load failure. +ITask* create(const BundleReader& bundle); + +} // namespace trtmc::gemma::edge_llm diff --git a/families/gemma/runtime/edge_llm/contract.h b/families/gemma/runtime/edge_llm/contract.h new file mode 100644 index 0000000000..5a76b7bb8f --- /dev/null +++ b/families/gemma/runtime/edge_llm/contract.h @@ -0,0 +1,57 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#pragma once + +#include "trtmc/task.h" + +#include +#include +#include +#include + +namespace trtmc::gemma::edge_llm { + +inline constexpr const char* kRevision = "e8b29522938901f6df19ebeedd4b69bc8edbcd97"; + +/// Return whether an artifact is a normalized file below one of the two Edge roots. +inline bool safe_artifact_path(const std::string& name) { + if (name.find('\\') != std::string::npos || name.find('\0') != std::string::npos) + return false; + const std::filesystem::path path(name); + if (path.is_absolute() || path.filename().empty()) + return false; + for (const auto& part : path) + if (part == "." || part == "..") + return false; + return path.generic_string() == name && + (name.rfind("edge_llm/engine/", 0) == 0 || name.rfind("edge_llm/checkpoint/", 0) == 0); +} + +/// Reject invalid sampling settings and controls with no equivalent Edge request API. +inline void validate_generation(const TextGenerationConfig& c, bool dspark = false) { + if (!dspark && c.temperature > 0 && c.top_k != 1) + throw std::invalid_argument("Gemma4 MTP supports only greedy generation"); + if (!std::isfinite(c.temperature) || c.temperature < 0 || !std::isfinite(c.top_p) || + c.top_p <= 0 || c.top_p > 1 || c.top_k < 0) + throw std::invalid_argument("Invalid Gemma4 Edge sampling parameters"); + if (c.min_p != 0 || c.seed != -1 || c.eos_token_id != -1 || c.repetition_penalty != 1 || + !c.lora_adapter_id.empty() || c.stop_on_boxed_answer || + (c.text_generation_mode != "auto" && c.text_generation_mode != "autoregressive") || + c.source_language_token_id != -1 || c.forced_bos_token_id != -1 || c.guidance_scale != -1 || + c.cfg_scale != -1 || c.num_steps != -1 || c.sde_gamma != -1 || !c.initial_latents.empty() || + !c.condition_latents.empty() || !c.condition_mask.empty() || !c.sampling_steps.empty() || + !c.sde_noises.empty() || c.block_length != 0 || c.confidence_threshold != -1) + throw std::invalid_argument("Requested generation controls are unsupported by Gemma4 Edge"); +} + +/// Enforce prompt and total capacity without allowing Edge to silently clip generation. +inline void validate_capacity(int prompt_tokens, int input_limit, int capacity, + std::int64_t generated_tokens) { + if (prompt_tokens <= 0 || prompt_tokens > input_limit || generated_tokens <= 0 || + generated_tokens > static_cast(capacity) - prompt_tokens) + throw std::invalid_argument("Gemma4 Edge prompt and generation exceed bundle capacity"); +} + +} // namespace trtmc::gemma::edge_llm diff --git a/families/gemma/runtime/edge_llm/device_link.cu b/families/gemma/runtime/edge_llm/device_link.cu new file mode 100644 index 0000000000..1ea4768d60 --- /dev/null +++ b/families/gemma/runtime/edge_llm/device_link.cu @@ -0,0 +1,5 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +// Enable the final CUDA device-link step for Edge's static runtime dependencies. diff --git a/families/gemma/runtime/edge_llm/request.h b/families/gemma/runtime/edge_llm/request.h new file mode 100644 index 0000000000..8df49ea8b4 --- /dev/null +++ b/families/gemma/runtime/edge_llm/request.h @@ -0,0 +1,68 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ +#pragma once + +#include "families/gemma/runtime/edge_llm/contract.h" + +#include +#include + +namespace trtmc::gemma::edge_llm { + +/// Match the checkpoint Jinja trim filter (Python Unicode whitespace), not the locale. +inline std::string trim_user_text(std::string_view text) { + constexpr std::string_view whitespace[]{ + "\t", "\n", "\v", "\f", "\r", + "\x1c", "\x1d", "\x1e", "\x1f", " ", + "\xc2\x85", "\xc2\xa0", "\xe1\x9a\x80", "\xe2\x80\x80", "\xe2\x80\x81", + "\xe2\x80\x82", "\xe2\x80\x83", "\xe2\x80\x84", "\xe2\x80\x85", "\xe2\x80\x86", + "\xe2\x80\x87", "\xe2\x80\x88", "\xe2\x80\x89", "\xe2\x80\x8a", "\xe2\x80\xa8", + "\xe2\x80\xa9", "\xe2\x80\xaf", "\xe2\x81\x9f", "\xe3\x80\x80"}; + for (;;) { + const auto previous = text.size(); + for (const auto space : whitespace) { + if (text.size() >= space.size() && text.substr(0, space.size()) == space) + text.remove_prefix(space.size()); + if (text.size() >= space.size() && text.substr(text.size() - space.size()) == space) + text.remove_suffix(space.size()); + } + if (text.size() == previous) + return std::string(text); + } +} + +/// Render the admitted Gemma4 checkpoint's single-user text template before Edge tokenization. +inline std::string render_chat(const std::string& prompt, bool thinking) { + std::string result = ""; + if (thinking) + result += "<|turn>system\n<|think|>\n\n"; + result += "<|turn>user\n" + trim_user_text(prompt) + "\n<|turn>model\n"; + if (!thinking) + result += "<|channel>thought\n"; + return result; +} + +/// Map Model Connect text arguments to the pinned Edge API; rejects unmapped controls. +inline trt_edgellm::rt::LLMGenerationRequest make_request(const std::string& prompt, + const TextGenerationConfig& config, + int default_length, bool dspark = false) { + validate_generation(config, dspark); + trt_edgellm::rt::LLMGenerationRequest request{}; + request.requests.resize(1); + const auto text = + config.use_chat_template ? render_chat(prompt, config.enable_thinking) : prompt; + request.requests.front().messages.push_back({"user", {{"text", text}}}); + // Edge 0.10.1's static Gemma template omits the disabled-thinking closure. + // Pass the faithful family-rendered prompt as raw text; never strip model output. + request.applyChatTemplate = false; + request.enableThinking = config.enable_thinking; + request.temperature = config.temperature; + request.topK = config.top_k; + request.topP = config.top_p; + request.maxGenerateLength = config.max_new_tokens > 0 ? config.max_new_tokens : default_length; + return request; +} + +} // namespace trtmc::gemma::edge_llm diff --git a/families/gemma/runtime/plugin.cpp b/families/gemma/runtime/plugin.cpp index 5c1dc213f1..c4ca9314fe 100644 --- a/families/gemma/runtime/plugin.cpp +++ b/families/gemma/runtime/plugin.cpp @@ -10,6 +10,9 @@ #include "families/gemma/runtime/plugin_helpers.h" #include "families/gemma/runtime/tensor_names.h" #include "trtmc/runtime/family_factory.h" +#ifdef TRTMC_HAS_EDGE_LLM +#include "families/gemma/runtime/edge_llm/adapter.h" +#endif #include #include @@ -207,6 +210,13 @@ DecoderModules load_modules(const FamilyContext& context, const RuntimeConfig& c } // namespace ITask* create(const FamilyContext& context) { + if (context.reader.find_section("edge_llm.json")) { +#ifdef TRTMC_HAS_EDGE_LLM + return edge_llm::create(context.reader); +#else + throw std::runtime_error("Gemma4 Edge bundle requires TRTMC_ENABLE_EDGELLM=ON"); +#endif + } const RuntimeConfig config = parse_runtime_config(context.reader); DecoderModules modules = load_modules(context, config); const cudaStream_t stream = modules.decode->stream(); diff --git a/families/gemma/support.py b/families/gemma/support.py index a5af06cf38..c63d37bc84 100644 --- a/families/gemma/support.py +++ b/families/gemma/support.py @@ -7,7 +7,7 @@ describe = family_support( - model_types=("gemma", "gemma2", "gemma3", "gemma3_text"), + model_types=("gemma", "gemma2", "gemma3", "gemma3_text", "gemma4_unified"), tasks=("text_generation",), default_task="text_generation", ) diff --git a/families/gemma/tests/cpp/test_gemma_sampler.cpp b/families/gemma/tests/cpp/test_gemma_sampler.cpp index 6b07604acb..22088e9c5e 100644 --- a/families/gemma/tests/cpp/test_gemma_sampler.cpp +++ b/families/gemma/tests/cpp/test_gemma_sampler.cpp @@ -5,6 +5,9 @@ #include "families/gemma/runtime/sampler.h" #include "trtmc/task.h" +#ifdef TRTMC_HAS_EDGE_LLM +#include "families/gemma/runtime/edge_llm/request.h" +#endif #include #include @@ -51,6 +54,39 @@ static void test_request_eos_overrides_model_defaults() { int main() { test_any_default_eos_stops_generation(); test_request_eos_overrides_model_defaults(); +#ifdef TRTMC_HAS_EDGE_LLM + trtmc::TextGenerationConfig config; + auto request = trtmc::gemma::edge_llm::make_request("hello", config, 128); + check(request.maxGenerateLength == 128 && request.topK == 1, + "MTP preserves the default greedy request"); + config.use_chat_template = true; + config.enable_thinking = false; + const auto chat = trtmc::gemma::edge_llm::make_request("\xe2\x80\x83hello \n", config, 128); + check(!chat.applyChatTemplate && + chat.requests.front().messages.front().contents.front().content == + "<|turn>user\nhello\n<|turn>model\n<|channel>thought\n", + "Gemma4 disabled thinking must match checkpoint prompt before tokenization"); + config.enable_thinking = true; + const auto thought = trtmc::gemma::edge_llm::make_request("hello", config, 128); + check(thought.requests.front().messages.front().contents.front().content == + "<|turn>system\n<|think|>\n\n<|turn>user\nhello\n<|turn>model\n", + "Gemma4 enabled thinking must match checkpoint system prefix"); + config.use_chat_template = false; + const auto raw = trtmc::gemma::edge_llm::make_request(" hello ", config, 128); + check(raw.requests.front().messages.front().contents.front().content == " hello ", + "Raw Gemma4 text must remain unmodified"); + config.top_k = 50; + bool rejected = false; + try { + (void)trtmc::gemma::edge_llm::make_request("hello", config, 128); + } catch (const std::invalid_argument&) { + rejected = true; + } + check(rejected, "MTP must not silently replace sampling with greedy"); + const auto sampled = trtmc::gemma::edge_llm::make_request("hello", config, 128, true); + check(sampled.topK == 50 && sampled.temperature == config.temperature, + "DSpark must preserve supported sampling controls"); +#endif if (failures > 0) { std::cerr << failures << " test(s) FAILED\n"; diff --git a/families/gemma/tests/test_e2e.py b/families/gemma/tests/test_e2e.py index 8d2e7ed223..328752b925 100644 --- a/families/gemma/tests/test_e2e.py +++ b/families/gemma/tests/test_e2e.py @@ -133,7 +133,7 @@ def _thresholds(case_name: str) -> dict[str, float]: return thresholds -def _build_bundle(manifest: dict, model_dir: Path, bundle: Path) -> None: +def _build_bundle(manifest: dict, model_dir: Path, bundle: Path, execution=None) -> None: quantization = manifest.get("quantization") assert quantization is None or isinstance(quantization, str) fp32_layers = tuple(manifest.get("fp32_layers", ())) @@ -148,7 +148,8 @@ def _build_bundle(manifest: dict, model_dir: Path, bundle: Path) -> None: tensor_parallel_size=manifest["tensor_parallel_size"], quantization=quantization, fp32_layers=fp32_layers, - ) + ), + execution=execution, ) assert bundle.is_file() and bundle.stat().st_size > 0, bundle diff --git a/families/gemma/tests/test_model_type_gate.py b/families/gemma/tests/test_model_type_gate.py index ce0bcfb354..f13e5d7293 100644 --- a/families/gemma/tests/test_model_type_gate.py +++ b/families/gemma/tests/test_model_type_gate.py @@ -68,12 +68,12 @@ def test_a_later_gemma_generation_is_refused(tmp_path: Path) -> None: """Gemma 3n and 4 need machinery this family still does not have. Gemma 4 adds vision and audio towers, per-layer input embeddings and - KV-shared layers; Gemma 3n is its own architecture again. Neither is built - here, so a prefix check would let them build a full-attention graph and + KV-shared layers; Gemma 3n is its own architecture again. Neither is built through the ordinary native entrypoint; + explicit paired Edge offload is separate, so a prefix check would let them build a full-attention graph and generate quietly wrong text. The refusal names the type so the message is actionable. Gemma 3 text is supported and is covered below. """ - for model_type in ("gemma3n", "gemma4", "gemma4_text"): + for model_type in ("gemma3n", "gemma4", "gemma4_text", "gemma4_unified"): directory = _model_dir(tmp_path / model_type.replace("_", ""), model_type) with pytest.raises(ValueError, match=re.escape(f"model_type={model_type!r}")): _build(directory) diff --git a/website/docs/features/model-families.md b/website/docs/features/model-families.md index cb3a6e18f6..d506d0a9b8 100644 --- a/website/docs/features/model-families.md +++ b/website/docs/features/model-families.md @@ -66,6 +66,23 @@ long-context qualification remain outside this contract. The committed-token-per-forward receipt is an algorithmic diagnostic, not a wall-clock speedup claim. +### Gemma4 paired ONNX execution + +The Gemma family also owns explicit Gemma4-12B target/assistant MTP and +Gemma4-12B/DSpark block7 execution through the optional pinned native Edge-LLM +SDK. These are text-only FP16 paired profiles, qualified on SM80; selecting a +Gemma4 checkpoint alone does not enable them or claim standalone native support. + +Provision the [native SDK](../user-guides/configure-runtime.md#optional-native-edge-llm-sdk), +then pass `--execution-variant mtp` or `--execution-variant dspark` with +`--companion draft=/path/to/checkpoint` to the build CLI. MTP is greedy-only; +DSpark preserves the supported sampling controls. Exact checkpoint revisions, +capacity bounds, validation results, unsupported controls and source-faithful +chat-template handling are documented in the +[owning Gemma recipe](https://github.com/NVIDIA/TensorRT-Model-Connect/blob/main/families/gemma/EDGE_LLM.md). +These local qualifications are separate from the registered manifest inventory +and do not imply that CI executes the paired cases. + ## Runtime and validation The directory name is also the runtime DSO identity: From 97cbe134c6d71b43c466f98f74cf4dc43cc7c9d4 Mon Sep 17 00:00:00 2001 From: Joshua Calafato Date: Tue, 22 Sep 2026 21:39:59 +0000 Subject: [PATCH 2/6] refactor(gemma): own Edge execution CLI and adapters Declare family-local CLI hooks and carry explicit paired inputs in a Gemma request. Route through the ordinary family builder, keep adapter code and CMake wiring in edge_llm directories, and preserve native behavior and all quality gates. Signed-off-by: Joshua Calafato --- .../gemma/{EDGE_LLM.md => edge_llm/README.md} | 35 +++- families/gemma/edge_llm/__init__.py | 4 + .../{edge_llm.py => edge_llm/builder.py} | 6 +- families/gemma/edge_llm/cli.py | 42 +++++ families/gemma/edge_llm/config.py | 88 +++++++++ families/gemma/model.py | 16 +- families/gemma/runtime/CMakeLists.txt | 17 +- families/gemma/runtime/edge_llm/Adapter.cmake | 19 ++ families/gemma/support.py | 1 + families/gemma/tests/test_e2e.py | 28 +-- families/gemma/tests/test_model_type_gate.py | 169 ++++++++++++++++++ website/docs/features/model-families.md | 2 +- 12 files changed, 384 insertions(+), 43 deletions(-) rename families/gemma/{EDGE_LLM.md => edge_llm/README.md} (74%) create mode 100644 families/gemma/edge_llm/__init__.py rename families/gemma/{edge_llm.py => edge_llm/builder.py} (98%) create mode 100644 families/gemma/edge_llm/cli.py create mode 100644 families/gemma/edge_llm/config.py create mode 100644 families/gemma/runtime/edge_llm/Adapter.cmake diff --git a/families/gemma/EDGE_LLM.md b/families/gemma/edge_llm/README.md similarity index 74% rename from families/gemma/EDGE_LLM.md rename to families/gemma/edge_llm/README.md index 91cc8dbf53..72d0e630e7 100644 --- a/families/gemma/EDGE_LLM.md +++ b/families/gemma/edge_llm/README.md @@ -7,15 +7,30 @@ paired path or claim native Gemma4 support. ## Build and inference -Use the [pinned native SDK provisioning](../../cmake/edgellm/README.md) with +Use the [pinned native SDK provisioning](../../../cmake/edge_llm/README.md) with `TRTMC_EDGELLM_ALL_KERNELS=ON` and `TRTMC_EDGELLM_ONNX=ON`, then configure the Model Connect runtime with `TRTMC_ENABLE_EDGELLM=ON`. Set `CMAKE_PREFIX_PATH` to the SDK installation. The source is official GitHub Edge-LLM 0.10.1, revision `e8b29522938901f6df19ebeedd4b69bc8edbcd97`; cross compilation is not used. -Pass `execution=BuildExecutionInputs("mtp", (NamedCheckpoint("draft", -draft_path),))` to the Python build API. The build CLI equivalent adds -`--execution-variant mtp --companion draft=/path/to/assistant`. +Use the existing build CLI with family-owned options (no checkpoint edits): + +```sh +trtmc build /path/to/target --family gemma --precision fp16 \ + -o /path/to/pair.bundle --execution-variant mtp \ + --companion draft=/path/to/assistant +``` + +`trtmc build /path/to/target --help` displays Gemma's options. Only Gemma +registers these flags; core does not interpret them or select Edge execution. +Python callers use `GemmaBuildRequest` and the family-owned +`BuildExecutionInputs`/`NamedCheckpoint` types from `families.gemma.edge_llm.config`, +then call the unchanged `tensorrt_model_connect.build(request)` API. +The existing ordinary Gemma `build(request, writer)` entrypoint chooses paired +Edge execution only for an explicit family request. Without it, native behavior +and unsupported-model rejection remain unchanged. Companion paths must name +existing local directories and are never inferred or downloaded. + The family forwards both unmodified checkpoints to the original Python ONNX exporter with `--mtp --mtp-draft-dir`, excluding image/audio branches for this text-only profile. The original native ONNX builder creates both speculative @@ -93,3 +108,15 @@ not claimed as byte-for-byte parity with the faulty upstream static template. MTP passes both unchanged quality gates with the same engine bundle after the request-only fix, including on the combined runtime. DSpark also passes both unchanged quality gates with the source-faithful prompt mapping. Initial failure evidence and the tokenizer audit are retained. + +## Family-owned CLI refactor validation + +The model qualification results above predate the CLI ownership refactor. +The refactor preserves exporter commands, engine composition, C++ inference +logic, reference outputs and quality thresholds. New coverage in the existing +family tests checks both CLI variants through ordinary core dispatch into the +family builder, malformed companions, typed request preservation and failed-build +bundle atomicity. Native adapter compilation and the sampler/pipeline C++ tests +were rerun successfully. Full checkpoint export/build/inference was not rerun +for this refactor; the successful qualification bundles were retired under the +approved artifact cleanup, so replay requires rebuilding those exact profiles. diff --git a/families/gemma/edge_llm/__init__.py b/families/gemma/edge_llm/__init__.py new file mode 100644 index 0000000000..2538fa21d9 --- /dev/null +++ b/families/gemma/edge_llm/__init__.py @@ -0,0 +1,4 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Self-contained Gemma Edge-LLM build integration.""" diff --git a/families/gemma/edge_llm.py b/families/gemma/edge_llm/builder.py similarity index 98% rename from families/gemma/edge_llm.py rename to families/gemma/edge_llm/builder.py index 8eb880d9e8..7fb014df07 100644 --- a/families/gemma/edge_llm.py +++ b/families/gemma/edge_llm/builder.py @@ -14,13 +14,15 @@ import tempfile import traceback -from tensorrt_model_connect.build import cmake_prefixes, detect_local_platform, subprocess_environment +from tensorrt_model_connect.build import ( + cmake_prefixes, detect_local_platform, subprocess_environment, +) EDGE_REVISION = "e8b29522938901f6df19ebeedd4b69bc8edbcd97" def local_target() -> dict: - """Return the executing worker identity supplied by generic build mechanics.""" + """Return the executing worker identity for this family-owned offload.""" return detect_local_platform() diff --git a/families/gemma/edge_llm/cli.py b/families/gemma/edge_llm/cli.py new file mode 100644 index 0000000000..fe3beb50ce --- /dev/null +++ b/families/gemma/edge_llm/cli.py @@ -0,0 +1,42 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Gemma-owned CLI options for explicit paired Edge execution.""" + +import argparse +from pathlib import Path + +from .config import BuildExecutionInputs, NamedCheckpoint, with_execution + + +def add_build_arguments(parser: argparse.ArgumentParser) -> None: + """Register options only when the Gemma family has been resolved.""" + parser.add_argument("--execution-variant", choices=("mtp", "dspark"), + help="Explicit Gemma paired execution mode") + parser.add_argument("--companion", action="append", default=[], metavar="ROLE=LOCAL_DIR", + help="Explicit local draft checkpoint; exactly one draft role is required") + + +def _execution_inputs(args: argparse.Namespace) -> BuildExecutionInputs | None: + """Parse only explicit local inputs; no variant list or model acquisition.""" + if args.command != "build": + return None + if args.execution_variant is None: + if args.companion: + raise ValueError("--companion requires --execution-variant") + return None + checkpoints = [] + for value in args.companion: + role, separator, directory = value.partition("=") + if not separator or not role or not directory: + raise ValueError("--companion must be ROLE=LOCAL_DIR") + if "://" in directory: + raise ValueError("--companion requires a local directory, not a URI") + checkpoints.append(NamedCheckpoint(role, Path(directory))) + return BuildExecutionInputs(args.execution_variant, tuple(checkpoints)) + + +def prepare_build_request(request, args): + """Attach a validated family-owned recipe before importing the GPU builder.""" + execution = _execution_inputs(args) + return request if execution is None else with_execution(request, execution) diff --git a/families/gemma/edge_llm/config.py b/families/gemma/edge_llm/config.py new file mode 100644 index 0000000000..50d5da83fe --- /dev/null +++ b/families/gemma/edge_llm/config.py @@ -0,0 +1,88 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Explicit Gemma paired-build inputs; no shared execution-mode contract.""" + +from dataclasses import dataclass, fields +from pathlib import Path +import re + +from tensorrt_model_connect.build import BuildRequest + + +_ID = re.compile(r"[a-z][a-z0-9_]*\Z") + + +def _validate_id(field: str, value: object) -> str: + if not isinstance(value, str) or _ID.fullmatch(value) is None: + raise ValueError(f"{field} must be a lowercase identifier containing only letters, digits, and underscores") + return value + + +@dataclass(frozen=True) +class NamedCheckpoint: + """One explicitly named local checkpoint; the family owns role semantics.""" + + role: str + model_dir: Path + + def __post_init__(self) -> None: + _validate_id("checkpoint role", self.role) + if not isinstance(self.model_dir, Path): + raise TypeError("checkpoint model_dir must be a Path") + if not self.model_dir.is_dir(): + raise ValueError(f"checkpoint must be an existing local directory: {self.model_dir}") + + +@dataclass(frozen=True) +class BuildExecutionInputs: + """Optional family-owned execution variant and immutable local companions. + + Gemma validates these inputs without fetching or inferring companions. + """ + + variant: str + checkpoints: tuple[NamedCheckpoint, ...] = () + + def __post_init__(self) -> None: + _validate_id("execution variant", self.variant) + if not isinstance(self.checkpoints, tuple) or any( + not isinstance(checkpoint, NamedCheckpoint) for checkpoint in self.checkpoints + ): + raise TypeError("checkpoints must be a tuple of NamedCheckpoint values") + roles = [checkpoint.role for checkpoint in self.checkpoints] + if len(roles) != len(set(roles)): + raise ValueError("checkpoint roles must be unique") + self.validate_local() + + def validate_local(self) -> None: + """Recheck local availability before dispatch, without acquiring inputs.""" + for checkpoint in self.checkpoints: + if not checkpoint.model_dir.is_dir(): + raise ValueError( + f"checkpoint must be an existing local directory: {checkpoint.model_dir}" + ) + + +@dataclass(frozen=True) +class GemmaBuildRequest(BuildRequest): + """Ordinary build inputs plus an explicitly requested Gemma execution recipe.""" + + execution: BuildExecutionInputs | None = None + + def __post_init__(self) -> None: + super().__post_init__() + if self.family != "gemma": + raise ValueError("GemmaBuildRequest requires the gemma family") + if self.execution is not None: + if not isinstance(self.execution, BuildExecutionInputs): + raise TypeError("execution must be BuildExecutionInputs") + self.execution.validate_local() + + +def with_execution(request: BuildRequest, execution: BuildExecutionInputs) -> GemmaBuildRequest: + """Preserve ordinary request fields and callback identity.""" + return GemmaBuildRequest( + **{field.name: getattr(request, field.name) for field in fields(BuildRequest)}, + execution=execution, + ) diff --git a/families/gemma/model.py b/families/gemma/model.py index 8418906039..2aab7248cc 100644 --- a/families/gemma/model.py +++ b/families/gemma/model.py @@ -395,6 +395,15 @@ def _decoder_engine_bytes(weights: "WeightDict", precision: str) -> int: def build(request: "BuildRequest", writer: "BundleWriter") -> None: """Build one Gemma bundle through family-owned code only.""" + from .edge_llm.config import GemmaBuildRequest + + if isinstance(request, GemmaBuildRequest) and request.execution is not None: + from .edge_llm.builder import build as build_pair + + request.execution.validate_local() + build_pair(request, writer, request.execution) + return + if request.dynamic_kv_cache: raise NotImplementedError("gemma does not support dynamic_kv_cache") @@ -532,10 +541,3 @@ def build(request: "BuildRequest", writer: "BundleWriter") -> None: path = model_dir / filename if path.is_file(): writer.add_bytes(filename, path.read_bytes()) - - -def build_with_inputs(request, writer, execution) -> None: - """Forward the explicit Gemma4 MTP or DSpark pair without changing standalone builds.""" - from .edge_llm import build as build_pair - - build_pair(request, writer, execution) diff --git a/families/gemma/runtime/CMakeLists.txt b/families/gemma/runtime/CMakeLists.txt index cc138d7716..9e7234f2a1 100644 --- a/families/gemma/runtime/CMakeLists.txt +++ b/families/gemma/runtime/CMakeLists.txt @@ -58,19 +58,4 @@ if(TRTMC_BUILD_TESTS) set_tests_properties(test_gemma_pipeline PROPERTIES SKIP_RETURN_CODE 77) endif() -# Optional complete-network offload; all model-specific orchestration stays here. -if(TARGET EdgeLLM::Core) - target_sources(trtmc_model_gemma PRIVATE edge_llm/adapter.cpp edge_llm/device_link.cu) - target_compile_definitions(trtmc_model_gemma PRIVATE TRTMC_HAS_EDGE_LLM=1) - target_link_libraries(trtmc_model_gemma PRIVATE EdgeLLM::Core) - set_target_properties(trtmc_model_gemma PROPERTIES - CUDA_ARCHITECTURES "${EdgeLLM_CUDA_ARCHITECTURE}" - CUDA_SEPARABLE_COMPILATION ON CUDA_RESOLVE_DEVICE_SYMBOLS ON) - add_custom_command(TARGET trtmc_model_gemma POST_BUILD - COMMAND ${CMAKE_COMMAND} -E copy_if_different - $ $ VERBATIM) - if(TRTMC_BUILD_TESTS) - target_compile_definitions(test_gemma_sampler PRIVATE TRTMC_HAS_EDGE_LLM=1) - target_link_libraries(test_gemma_sampler PRIVATE EdgeLLM::Core) - endif() -endif() +include("${CMAKE_CURRENT_LIST_DIR}/edge_llm/Adapter.cmake") diff --git a/families/gemma/runtime/edge_llm/Adapter.cmake b/families/gemma/runtime/edge_llm/Adapter.cmake new file mode 100644 index 0000000000..04e8427575 --- /dev/null +++ b/families/gemma/runtime/edge_llm/Adapter.cmake @@ -0,0 +1,19 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +# Optional complete-network offload; all model-specific orchestration stays here. +if(TARGET EdgeLLM::Core) + target_sources(trtmc_model_gemma PRIVATE "${CMAKE_CURRENT_LIST_DIR}/adapter.cpp" "${CMAKE_CURRENT_LIST_DIR}/device_link.cu") + target_compile_definitions(trtmc_model_gemma PRIVATE TRTMC_HAS_EDGE_LLM=1) + target_link_libraries(trtmc_model_gemma PRIVATE EdgeLLM::Core) + set_target_properties(trtmc_model_gemma PROPERTIES + CUDA_ARCHITECTURES "${EdgeLLM_CUDA_ARCHITECTURE}" + CUDA_SEPARABLE_COMPILATION ON CUDA_RESOLVE_DEVICE_SYMBOLS ON) + add_custom_command(TARGET trtmc_model_gemma POST_BUILD + COMMAND ${CMAKE_COMMAND} -E copy_if_different + $ $ VERBATIM) + if(TRTMC_BUILD_TESTS) + target_compile_definitions(test_gemma_sampler PRIVATE TRTMC_HAS_EDGE_LLM=1) + target_link_libraries(test_gemma_sampler PRIVATE EdgeLLM::Core) + endif() +endif() diff --git a/families/gemma/support.py b/families/gemma/support.py index c63d37bc84..fa8399407e 100644 --- a/families/gemma/support.py +++ b/families/gemma/support.py @@ -10,4 +10,5 @@ model_types=("gemma", "gemma2", "gemma3", "gemma3_text", "gemma4_unified"), tasks=("text_generation",), default_task="text_generation", + build_cli_module="edge_llm.cli", ) diff --git a/families/gemma/tests/test_e2e.py b/families/gemma/tests/test_e2e.py index 328752b925..aa338e9ee0 100644 --- a/families/gemma/tests/test_e2e.py +++ b/families/gemma/tests/test_e2e.py @@ -137,20 +137,22 @@ def _build_bundle(manifest: dict, model_dir: Path, bundle: Path, execution=None) quantization = manifest.get("quantization") assert quantization is None or isinstance(quantization, str) fp32_layers = tuple(manifest.get("fp32_layers", ())) - build( - BuildRequest( - model_dir=model_dir, - output_path=bundle, - family=_FAMILY, - task="text_generation", - precision=manifest["precision"], - max_sequence_length=manifest["max_sequence_length"], - tensor_parallel_size=manifest["tensor_parallel_size"], - quantization=quantization, - fp32_layers=fp32_layers, - ), - execution=execution, + request = BuildRequest( + model_dir=model_dir, + output_path=bundle, + family=_FAMILY, + task="text_generation", + precision=manifest["precision"], + max_sequence_length=manifest["max_sequence_length"], + tensor_parallel_size=manifest["tensor_parallel_size"], + quantization=quantization, + fp32_layers=fp32_layers, ) + if execution is not None: + from families.gemma.edge_llm.config import with_execution + + request = with_execution(request, execution) + build(request) assert bundle.is_file() and bundle.stat().st_size > 0, bundle diff --git a/families/gemma/tests/test_model_type_gate.py b/families/gemma/tests/test_model_type_gate.py index f13e5d7293..b945b8e391 100644 --- a/families/gemma/tests/test_model_type_gate.py +++ b/families/gemma/tests/test_model_type_gate.py @@ -5,7 +5,9 @@ from __future__ import annotations +import importlib import json +from dataclasses import FrozenInstanceError import re from pathlib import Path @@ -22,6 +24,13 @@ pytest.skip("tensorrt_model_connect requires TensorRT", allow_module_level=True) +from families.gemma.edge_llm.config import ( + BuildExecutionInputs, GemmaBuildRequest, NamedCheckpoint, with_execution, +) + +build_core = importlib.import_module("tensorrt_model_connect.build") + + def _model_dir(tmp_path: Path, model_type: str) -> Path: tmp_path.mkdir(parents=True, exist_ok=True) (tmp_path / "config.json").write_text( @@ -132,3 +141,163 @@ def _build_with(model_dir: Path, *, precision: str) -> None: ), writer=None, ) + + +def execution_request(root: Path) -> BuildRequest: + return BuildRequest(root, root / "model.bundle", "gemma", "text_generation", "fp16") + + +def inputs(root: Path) -> BuildExecutionInputs: + return BuildExecutionInputs("mtp", (NamedCheckpoint("draft", root),)) + + +def test_execution_inputs_are_immutable(tmp_path): + execution = inputs(tmp_path) + with pytest.raises(FrozenInstanceError): + execution.variant = "other" + with pytest.raises(FrozenInstanceError): + execution.checkpoints[0].role = "other" + assert execution.checkpoints[0].model_dir is tmp_path + + +@pytest.mark.parametrize("value", ["", "../bad", "UPPER", "a-b", "a.b"]) +def test_invalid_role_and_variant(tmp_path, value): + with pytest.raises(ValueError, match="lowercase identifier"): + NamedCheckpoint(value, tmp_path) + with pytest.raises(ValueError, match="lowercase identifier"): + BuildExecutionInputs(value) + + +def test_execution_requires_immutable_typed_companions(tmp_path): + checkpoint = NamedCheckpoint("draft", tmp_path) + with pytest.raises(TypeError, match="tuple"): + BuildExecutionInputs("mtp", [checkpoint]) + with pytest.raises(TypeError, match="NamedCheckpoint"): + BuildExecutionInputs("mtp", (object(),)) + with pytest.raises(ValueError, match="unique"): + BuildExecutionInputs("mtp", (checkpoint, checkpoint)) + with pytest.raises(TypeError, match="Path"): + NamedCheckpoint("draft", str(tmp_path)) + + +def test_local_checkpoint_required_and_rechecked(tmp_path, monkeypatch): + with pytest.raises(ValueError, match="existing local directory"): + NamedCheckpoint("draft", tmp_path / "missing") + file = tmp_path / "file" + file.write_text("not a directory") + with pytest.raises(ValueError, match="existing local directory"): + NamedCheckpoint("draft", file) + directory = tmp_path / "companion" + directory.mkdir() + execution = inputs(directory) + directory.rmdir() + monkeypatch.setattr(build_core, "_select_backend", lambda _: pytest.fail("backend touched")) + with pytest.raises(ValueError, match="existing local directory"): + with_execution(execution_request(tmp_path), execution) + + +def test_untyped_execution_fails_before_side_effects(tmp_path, monkeypatch): + monkeypatch.setattr(build_core, "_select_backend", lambda _: pytest.fail("backend touched")) + with pytest.raises(TypeError, match="BuildExecutionInputs"): + with_execution(execution_request(tmp_path), {"variant": "paired"}) + + + + +@pytest.mark.parametrize("variant", ["mtp", "dspark"]) +def test_edge_cli_routes_through_the_ordinary_family_entrypoint(tmp_path, monkeypatch, variant): + from families.gemma.edge_llm import builder as edge_builder + from tensorrt_model_connect import build_cli + + source = _model_dir(tmp_path / "target", "gemma4_unified") + draft = tmp_path / "draft" + draft.mkdir() + output = tmp_path / "pair.bundle" + seen = [] + + def paired(request, writer, execution): + assert isinstance(request, GemmaBuildRequest) + assert request.execution is execution + seen.append(execution) + writer.set_header(family="gemma", task=request.task, backend=request.backend) + writer.add_json("edge-test.json", {"variant": execution.variant}) + + monkeypatch.setattr(edge_builder, "build", paired) + assert build_cli.main([ + "build", str(source), "--family", "gemma", "--precision", "fp16", + "-o", str(output), "--execution-variant", variant, + "--companion", f"draft={draft}", + ]) == 0 + assert seen == [BuildExecutionInputs(variant, (NamedCheckpoint("draft", draft),))] + assert output.is_file() + + +@pytest.mark.parametrize("options", [ + ["--companion", "draft=/missing"], + ["--execution-variant", "mtp", "--companion", "missing_separator"], + ["--execution-variant", "mtp", "--companion", "=path"], + ["--execution-variant", "mtp", "--companion", "draft="], + ["--execution-variant", "mtp", "--companion", "draft=https://example.com/model"], +]) +def test_bad_edge_cli_inputs_fail_before_backend_and_bundle(tmp_path, monkeypatch, options): + from tensorrt_model_connect import build_cli + + source = _model_dir(tmp_path / "target", "gemma4_unified") + output = tmp_path / "pair.bundle" + monkeypatch.setattr(build_core, "_select_backend", lambda *_: pytest.fail("backend touched")) + monkeypatch.setattr(build_core, "BundleWriter", lambda *_: pytest.fail("writer created")) + with pytest.raises(ValueError): + build_cli.main(["build", str(source), "--family", "gemma", "-o", str(output), *options]) + assert not output.exists() + + +@pytest.mark.parametrize("failure", [RuntimeError("paired build failed"), KeyboardInterrupt()]) +def test_edge_family_failure_preserves_existing_publication(tmp_path, monkeypatch, failure): + from families.gemma.edge_llm import builder as edge_builder + + request = with_execution(execution_request(tmp_path), inputs(tmp_path)) + request.output_path.write_bytes(b"previous valid publication") + + def fail(actual, writer, execution): + writer.set_header(family="gemma", task=actual.task, backend=actual.backend) + writer.add_json("edge-test.json", {"variant": execution.variant}) + raise failure + + monkeypatch.setattr(edge_builder, "build", fail) + with pytest.raises(type(failure)) as caught: + build_core.build(request) + assert caught.value is failure + assert request.output_path.read_bytes() == b"previous valid publication" + assert sorted(path.name for path in tmp_path.iterdir()) == ["model.bundle"] + + +def test_edge_request_keeps_graph_callback_and_all_ordinary_fields(tmp_path): + from dataclasses import fields, replace + + def callback(layer): + return layer + ordinary = replace(execution_request(tmp_path), graph_transform=callback, max_sequence_length=128) + extended = with_execution(ordinary, inputs(tmp_path)) + for field in fields(BuildRequest): + assert getattr(extended, field.name) is getattr(ordinary, field.name) + with pytest.raises(FrozenInstanceError): + extended.execution = None + + +def test_family_cli_without_extra_options_preserves_native_request(tmp_path): + import argparse + from families.gemma.edge_llm import cli + + parser = argparse.ArgumentParser() + parser.set_defaults(command="build") + cli.add_build_arguments(parser) + request = execution_request(tmp_path) + assert cli.prepare_build_request(request, parser.parse_args([])) is request + + +def test_paired_request_cannot_be_dispatched_to_another_family(tmp_path): + from dataclasses import replace + + request = replace(execution_request(tmp_path), family="another_owner") + with pytest.raises(ValueError, match="requires the gemma family"): + with_execution(request, inputs(tmp_path)) diff --git a/website/docs/features/model-families.md b/website/docs/features/model-families.md index d506d0a9b8..b9173b51fd 100644 --- a/website/docs/features/model-families.md +++ b/website/docs/features/model-families.md @@ -79,7 +79,7 @@ then pass `--execution-variant mtp` or `--execution-variant dspark` with DSpark preserves the supported sampling controls. Exact checkpoint revisions, capacity bounds, validation results, unsupported controls and source-faithful chat-template handling are documented in the -[owning Gemma recipe](https://github.com/NVIDIA/TensorRT-Model-Connect/blob/main/families/gemma/EDGE_LLM.md). +[owning Gemma recipe](https://github.com/NVIDIA/TensorRT-Model-Connect/blob/main/families/gemma/edge_llm/README.md). These local qualifications are separate from the registered manifest inventory and do not imply that CI executes the paired cases. From d34000f8fa9e3c353f03b7af5f6fd6df7a96eb63 Mon Sep 17 00:00:00 2001 From: Joshua Calafato Date: Tue, 22 Sep 2026 22:12:59 +0000 Subject: [PATCH 3/6] fix(gemma): reject incomplete Edge SDK targets Signed-off-by: Joshua Calafato --- families/gemma/runtime/edge_llm/Adapter.cmake | 3 +++ 1 file changed, 3 insertions(+) diff --git a/families/gemma/runtime/edge_llm/Adapter.cmake b/families/gemma/runtime/edge_llm/Adapter.cmake index 04e8427575..8f962b25c5 100644 --- a/families/gemma/runtime/edge_llm/Adapter.cmake +++ b/families/gemma/runtime/edge_llm/Adapter.cmake @@ -3,6 +3,9 @@ # Optional complete-network offload; all model-specific orchestration stays here. if(TARGET EdgeLLM::Core) + if(NOT TARGET EdgeLLM::Plugin) + message(FATAL_ERROR "Gemma Edge adapter requires the complete EdgeLLM package (Core and Plugin)") + endif() target_sources(trtmc_model_gemma PRIVATE "${CMAKE_CURRENT_LIST_DIR}/adapter.cpp" "${CMAKE_CURRENT_LIST_DIR}/device_link.cu") target_compile_definitions(trtmc_model_gemma PRIVATE TRTMC_HAS_EDGE_LLM=1) target_link_libraries(trtmc_model_gemma PRIVATE EdgeLLM::Core) From 01015c18943e0e5727af5a3262cc9b8694a32006 Mon Sep 17 00:00:00 2001 From: Joshua Calafato Date: Wed, 23 Sep 2026 19:46:54 +0000 Subject: [PATCH 4/6] refactor(gemma): use declared family CLI Reuse the existing family CLI protocol instead of extending the shared parser. Own the command description, request contract and build lifecycle; keep Edge companion semantics inside this family. Preserve legacy callers through strict conversion and keep numerical acceptance gates unchanged. Signed-off-by: Joshua Calafato --- families/gemma/build_request.py | 96 +++++++++++++++ families/gemma/cli.json | 122 +++++++++++++++++++ families/gemma/cli.py | 57 +++++++++ families/gemma/edge_llm/README.md | 14 ++- families/gemma/edge_llm/cli.py | 33 ++--- families/gemma/edge_llm/config.py | 5 +- families/gemma/model.py | 5 +- families/gemma/support.py | 1 - families/gemma/tests/test_model_type_gate.py | 75 ++++++++++-- website/docs/features/model-families.md | 5 + 10 files changed, 371 insertions(+), 42 deletions(-) create mode 100644 families/gemma/build_request.py create mode 100644 families/gemma/cli.json create mode 100644 families/gemma/cli.py diff --git a/families/gemma/build_request.py b/families/gemma/build_request.py new file mode 100644 index 0000000000..712450a0e6 --- /dev/null +++ b/families/gemma/build_request.py @@ -0,0 +1,96 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""gemma build inputs and strict compatibility for existing Python callers.""" + +from __future__ import annotations + +from dataclasses import dataclass, fields +from pathlib import Path +import re +from typing import ClassVar + +from tensorrt_model_connect.graph_transform import GraphTransform + + +def _validate_id(field: str, value: str) -> None: + if not isinstance(value, str) or re.fullmatch(r"[a-z][a-z0-9_]*", value) is None: + raise ValueError(f"{field} must be a lowercase identifier") + + +@dataclass(frozen=True) +class BuildRequest: + """gemma-owned inputs; unsupported legacy controls are read-only defaults.""" + + model_dir: Path + output_path: Path + family: str + task: str + precision: str + backend: str = "trt" + max_sequence_length: int | None = None + image_height: ClassVar[int | None] = None + image_width: ClassVar[int | None] = None + video_num_frames: ClassVar[int | None] = None + max_batch_size: ClassVar[int] = 1 + tensor_parallel_size: int = 1 + context_parallel_size: ClassVar[int] = 1 + quantization: ClassVar[str | None] = None + fp32_layers: ClassVar[tuple[int, ...]] = () + dynamic_kv_cache: ClassVar[bool] = False + verbose: bool = False + graph_transform: GraphTransform | None = None + + def __post_init__(self) -> None: + if not self.precision: + raise ValueError("precision must be non-empty") + _validate_id("family", self.family) + _validate_id("task", self.task) + if self.backend not in {"trt", "trt_rtx"}: + raise ValueError("backend must be 'trt' or 'trt_rtx'") + if self.max_sequence_length is not None and self.max_sequence_length < 1: + raise ValueError("max_sequence_length must be positive") + for field in ("image_height", "image_width", "video_num_frames"): + value = getattr(self, field) + if value is not None and value < 1: + raise ValueError(f"{field} must be positive") + if self.max_batch_size < 1: + raise ValueError("max_batch_size must be positive") + if self.tensor_parallel_size < 1: + raise ValueError("tensor_parallel_size must be positive") + if self.context_parallel_size < 1: + raise ValueError("context_parallel_size must be positive") + if self.quantization is not None and not self.quantization: + raise ValueError("quantization must be non-empty when provided") + if any(layer < 0 for layer in self.fp32_layers): + raise ValueError("fp32_layers must contain non-negative indices") + if not isinstance(self.dynamic_kv_cache, bool): + raise ValueError("dynamic_kv_cache must be a bool") + if self.graph_transform is not None and not callable(self.graph_transform): + raise ValueError("graph_transform must be callable when provided") + + +def coerce_request(request: object) -> BuildRequest: + """Reject unsupported/unknown legacy inputs before converting to owner fields.""" + if isinstance(request, BuildRequest): + return request + unsupported = { + "image_height": None, + "image_width": None, + "video_num_frames": None, + "max_batch_size": 1, + "context_parallel_size": 1, + "quantization": None, + "fp32_layers": (), + "dynamic_kv_cache": False, + } + for name, default in unsupported.items(): + value = getattr(request, name, default) + if name == "quantization" and value == "none": + continue + if value != default: + raise NotImplementedError(f"gemma does not support {name}") + names = {field.name for field in fields(BuildRequest)} + if unknown := set(vars(request)) - names - set(unsupported): + raise ValueError(f"unknown gemma build inputs: {sorted(unknown)}") + return BuildRequest(**{name: getattr(request, name) for name in names}) diff --git a/families/gemma/cli.json b/families/gemma/cli.json new file mode 100644 index 0000000000..be801c5e83 --- /dev/null +++ b/families/gemma/cli.json @@ -0,0 +1,122 @@ +{ + "version": 1, + "commands": [ + { + "name": "build", + "help": "Build one gemma TensorRT bundle", + "executor": "python", + "handler": "cli:build", + "arguments": [ + { + "name": "model", + "type": "string", + "help": "Hugging Face model ID or local snapshot" + }, + { + "name": "output", + "flags": [ + "-o", + "--output" + ], + "type": "path", + "required": true + }, + { + "name": "revision", + "flags": [ + "--revision" + ], + "type": "string" + }, + { + "name": "task", + "flags": [ + "--task" + ], + "type": "string", + "choices": [ + "text_generation" + ], + "default": "text_generation" + }, + { + "name": "precision", + "flags": [ + "--precision" + ], + "type": "string", + "choices": [ + "fp32", + "fp16", + "bf16" + ], + "default": "fp32" + }, + { + "name": "backend", + "flags": [ + "--backend" + ], + "type": "string", + "choices": [ + "trt", + "trt_rtx" + ], + "default": "trt" + }, + { + "name": "max_sequence_length", + "flags": [ + "--max-sequence-length" + ], + "type": "int" + }, + { + "name": "tensor_parallel_size", + "flags": [ + "--tensor-parallel-size" + ], + "type": "int", + "choices": [ + 1, + 2, + 4, + 8 + ], + "default": 1 + }, + { + "name": "verbose", + "flags": [ + "--verbose" + ], + "type": "bool", + "action": "store_true", + "default": false + }, + { + "name": "execution_variant", + "flags": [ + "--execution-variant" + ], + "type": "string", + "choices": [ + "mtp", + "dspark" + ], + "help": "Explicit paired execution mode; no automatic draft discovery" + }, + { + "name": "companion", + "flags": [ + "--companion" + ], + "type": "string", + "action": "append", + "default": [], + "help": "Local companion checkpoint as ROLE=LOCAL_DIR" + } + ] + } + ] +} diff --git a/families/gemma/cli.py b/families/gemma/cli.py new file mode 100644 index 0000000000..efcd0fa2cb --- /dev/null +++ b/families/gemma/cli.py @@ -0,0 +1,57 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""gemma-owned build command and typed inputs; importing this module is CPU-only.""" + +from __future__ import annotations + +from dataclasses import replace +from pathlib import Path + +from tensorrt_model_connect.build import select_backend +from tensorrt_model_connect.bundle_writer import BundleWriter +from tensorrt_model_connect.graph_transform import graph_transform +from tensorrt_model_connect.model_support import load_model_metadata, resolve_family, resolve_model + +from .build_request import BuildRequest + +from .edge_llm.cli import execution_inputs +from .edge_llm.config import with_execution + + +def build_bundle(request: BuildRequest, output: Path) -> None: + """Build and publish through the owning model, preserving atomic failure.""" + request = replace(request, output_path=output) + select_backend(request.backend) + from .model import build as build_model + + writer = BundleWriter(output) + try: + with graph_transform(request.graph_transform): + build_model(request, writer) + writer.finish() + except BaseException: + writer.abort() + raise + +def build( + *, model: str, output: Path, revision: str | None = None, + task: str = "text_generation", precision: str = "fp32", backend: str = "trt", + max_sequence_length: int | None = None, tensor_parallel_size: int = 1, + verbose: bool = False, + execution_variant: str | None = None, companion: list[str] | tuple[str, ...] = (), +) -> int: + """Run the declared owner command; help never imports this handler.""" + execution = execution_inputs(execution_variant, companion) + model_dir = resolve_model(model, revision) + resolve_family(load_model_metadata(model_dir), "gemma") + request = BuildRequest( + model_dir=model_dir, output_path=output, family="gemma", + task=task, precision=precision, backend=backend, + max_sequence_length=max_sequence_length, tensor_parallel_size=tensor_parallel_size, + verbose=verbose, + ) + if execution is not None: + request = with_execution(request, execution) + build_bundle(request, output) + return 0 diff --git a/families/gemma/edge_llm/README.md b/families/gemma/edge_llm/README.md index 72d0e630e7..87c85468b5 100644 --- a/families/gemma/edge_llm/README.md +++ b/families/gemma/edge_llm/README.md @@ -13,16 +13,16 @@ Model Connect runtime with `TRTMC_ENABLE_EDGELLM=ON`. Set `CMAKE_PREFIX_PATH` to the SDK installation. The source is official GitHub Edge-LLM 0.10.1, revision `e8b29522938901f6df19ebeedd4b69bc8edbcd97`; cross compilation is not used. -Use the existing build CLI with family-owned options (no checkpoint edits): +Use the existing family CLI protocol with family-owned options (no checkpoint edits): ```sh -trtmc build /path/to/target --family gemma --precision fp16 \ +trtmc gemma build /path/to/target --precision fp16 \ -o /path/to/pair.bundle --execution-variant mtp \ --companion draft=/path/to/assistant ``` -`trtmc build /path/to/target --help` displays Gemma's options. Only Gemma -registers these flags; core does not interpret them or select Edge execution. +`trtmc gemma build /path/to/target --help` displays Gemma's options. Only Gemma +declares these flags in cli.json; core does not interpret them or select Edge execution. Python callers use `GemmaBuildRequest` and the family-owned `BuildExecutionInputs`/`NamedCheckpoint` types from `families.gemma.edge_llm.config`, then call the unchanged `tensorrt_model_connect.build(request)` API. @@ -120,3 +120,9 @@ bundle atomicity. Native adapter compilation and the sampler/pipeline C++ tests were rerun successfully. Full checkpoint export/build/inference was not rerun for this refactor; the successful qualification bundles were retired under the approved artifact cleanup, so replay requires rebuilding those exact profiles. + +## Declared build command + +This family uses the existing cli.json protocol introduced in #1310. The family\nowns its declaration, typed inputs and Python handler. The handler adapts those\ninputs to the unchanged builder API, preserving native/Edge dispatch and bundle\npublication. The legacy flat build command remains available for its existing\nordinary options; new family options use `trtmc gemma build`. +Help is offline and does not need a local checkpoint. No shared parser hook or +family registry entry is added. diff --git a/families/gemma/edge_llm/cli.py b/families/gemma/edge_llm/cli.py index fe3beb50ce..467d8188d8 100644 --- a/families/gemma/edge_llm/cli.py +++ b/families/gemma/edge_llm/cli.py @@ -3,40 +3,27 @@ """Gemma-owned CLI options for explicit paired Edge execution.""" -import argparse from pathlib import Path -from .config import BuildExecutionInputs, NamedCheckpoint, with_execution +from .config import BuildExecutionInputs, NamedCheckpoint -def add_build_arguments(parser: argparse.ArgumentParser) -> None: - """Register options only when the Gemma family has been resolved.""" - parser.add_argument("--execution-variant", choices=("mtp", "dspark"), - help="Explicit Gemma paired execution mode") - parser.add_argument("--companion", action="append", default=[], metavar="ROLE=LOCAL_DIR", - help="Explicit local draft checkpoint; exactly one draft role is required") - - -def _execution_inputs(args: argparse.Namespace) -> BuildExecutionInputs | None: +def execution_inputs( + execution_variant: str | None, companion: list[str] | tuple[str, ...] = (), +) -> BuildExecutionInputs | None: """Parse only explicit local inputs; no variant list or model acquisition.""" - if args.command != "build": - return None - if args.execution_variant is None: - if args.companion: + if execution_variant is None: + if companion: raise ValueError("--companion requires --execution-variant") return None + if execution_variant not in ['mtp','dspark']: + raise ValueError("unsupported gemma execution variant") checkpoints = [] - for value in args.companion: + for value in companion: role, separator, directory = value.partition("=") if not separator or not role or not directory: raise ValueError("--companion must be ROLE=LOCAL_DIR") if "://" in directory: raise ValueError("--companion requires a local directory, not a URI") checkpoints.append(NamedCheckpoint(role, Path(directory))) - return BuildExecutionInputs(args.execution_variant, tuple(checkpoints)) - - -def prepare_build_request(request, args): - """Attach a validated family-owned recipe before importing the GPU builder.""" - execution = _execution_inputs(args) - return request if execution is None else with_execution(request, execution) + return BuildExecutionInputs(execution_variant, tuple(checkpoints)) diff --git a/families/gemma/edge_llm/config.py b/families/gemma/edge_llm/config.py index 50d5da83fe..98a35b21fb 100644 --- a/families/gemma/edge_llm/config.py +++ b/families/gemma/edge_llm/config.py @@ -7,7 +7,7 @@ from pathlib import Path import re -from tensorrt_model_connect.build import BuildRequest +from ..build_request import BuildRequest, coerce_request _ID = re.compile(r"[a-z][a-z0-9_]*\Z") @@ -81,7 +81,8 @@ def __post_init__(self) -> None: def with_execution(request: BuildRequest, execution: BuildExecutionInputs) -> GemmaBuildRequest: - """Preserve ordinary request fields and callback identity.""" + """Preserve supported ordinary request fields and callback identity.""" + request = coerce_request(request) return GemmaBuildRequest( **{field.name: getattr(request, field.name) for field in fields(BuildRequest)}, execution=execution, diff --git a/families/gemma/model.py b/families/gemma/model.py index 2aab7248cc..df30775772 100644 --- a/families/gemma/model.py +++ b/families/gemma/model.py @@ -27,7 +27,7 @@ if TYPE_CHECKING: - from tensorrt_model_connect.build import BuildRequest + from .build_request import BuildRequest from tensorrt_model_connect.bundle_writer import BundleWriter @@ -395,6 +395,9 @@ def _decoder_engine_bytes(weights: "WeightDict", precision: str) -> int: def build(request: "BuildRequest", writer: "BundleWriter") -> None: """Build one Gemma bundle through family-owned code only.""" + from .build_request import coerce_request + + request = coerce_request(request) from .edge_llm.config import GemmaBuildRequest if isinstance(request, GemmaBuildRequest) and request.execution is not None: diff --git a/families/gemma/support.py b/families/gemma/support.py index fa8399407e..c63d37bc84 100644 --- a/families/gemma/support.py +++ b/families/gemma/support.py @@ -10,5 +10,4 @@ model_types=("gemma", "gemma2", "gemma3", "gemma3_text", "gemma4_unified"), tasks=("text_generation",), default_task="text_generation", - build_cli_module="edge_llm.cli", ) diff --git a/families/gemma/tests/test_model_type_gate.py b/families/gemma/tests/test_model_type_gate.py index b945b8e391..ae1f950266 100644 --- a/families/gemma/tests/test_model_type_gate.py +++ b/families/gemma/tests/test_model_type_gate.py @@ -207,7 +207,7 @@ def test_untyped_execution_fails_before_side_effects(tmp_path, monkeypatch): @pytest.mark.parametrize("variant", ["mtp", "dspark"]) def test_edge_cli_routes_through_the_ordinary_family_entrypoint(tmp_path, monkeypatch, variant): from families.gemma.edge_llm import builder as edge_builder - from tensorrt_model_connect import build_cli + from tensorrt_model_connect import family_cli as build_cli source = _model_dir(tmp_path / "target", "gemma4_unified") draft = tmp_path / "draft" @@ -223,8 +223,8 @@ def paired(request, writer, execution): writer.add_json("edge-test.json", {"variant": execution.variant}) monkeypatch.setattr(edge_builder, "build", paired) - assert build_cli.main([ - "build", str(source), "--family", "gemma", "--precision", "fp16", + assert build_cli.main(["gemma", + "build", str(source), "--precision", "fp16", "-o", str(output), "--execution-variant", variant, "--companion", f"draft={draft}", ]) == 0 @@ -240,14 +240,14 @@ def paired(request, writer, execution): ["--execution-variant", "mtp", "--companion", "draft=https://example.com/model"], ]) def test_bad_edge_cli_inputs_fail_before_backend_and_bundle(tmp_path, monkeypatch, options): - from tensorrt_model_connect import build_cli + from tensorrt_model_connect import family_cli as build_cli source = _model_dir(tmp_path / "target", "gemma4_unified") output = tmp_path / "pair.bundle" monkeypatch.setattr(build_core, "_select_backend", lambda *_: pytest.fail("backend touched")) monkeypatch.setattr(build_core, "BundleWriter", lambda *_: pytest.fail("writer created")) with pytest.raises(ValueError): - build_cli.main(["build", str(source), "--family", "gemma", "-o", str(output), *options]) + build_cli.main(["gemma", "build", str(source), "-o", str(output), *options]) assert not output.exists() @@ -285,14 +285,9 @@ def callback(layer): def test_family_cli_without_extra_options_preserves_native_request(tmp_path): - import argparse from families.gemma.edge_llm import cli - parser = argparse.ArgumentParser() - parser.set_defaults(command="build") - cli.add_build_arguments(parser) - request = execution_request(tmp_path) - assert cli.prepare_build_request(request, parser.parse_args([])) is request + assert cli.execution_inputs(None) is None def test_paired_request_cannot_be_dispatched_to_another_family(tmp_path): @@ -301,3 +296,61 @@ def test_paired_request_cannot_be_dispatched_to_another_family(tmp_path): request = replace(execution_request(tmp_path), family="another_owner") with pytest.raises(ValueError, match="requires the gemma family"): with_execution(request, inputs(tmp_path)) + + +@pytest.mark.parametrize("options", [[], ["--precision", "fp16", "--max-sequence-length", "64"]]) +def test_declared_build_matches_legacy_request(tmp_path, monkeypatch, options): + """Owner command preserves ordinary request defaults and explicit controls.""" + import json + from families.gemma import cli as owner + from tensorrt_model_connect import build_cli, family_cli + + source = tmp_path / "checkpoint" + source.mkdir() + (source / "config.json").write_text(json.dumps({"model_type": "gemma"})) + output = tmp_path / "model.bundle" + captured = [] + monkeypatch.setattr(owner, "build_bundle", lambda request, output: captured.append(request)) + monkeypatch.setattr(build_cli, "build", captured.append) + args = [str(source), "-o", str(output), *options] + assert family_cli.main(["gemma", "build", *args]) == 0 + assert build_cli.main(["build", *args, "--family", "gemma"]) == 0 + assert len(captured) == 2 + from dataclasses import fields + assert isinstance(captured[0], owner.BuildRequest) + for field in fields(captured[1]): + assert getattr(captured[0], field.name) == getattr(captured[1], field.name) + from dataclasses import replace + from families.gemma.build_request import coerce_request + assert coerce_request(captured[1]) == captured[0] + with pytest.raises(NotImplementedError, match="image_height"): + coerce_request(replace(captured[1], image_height=32)) + from types import SimpleNamespace + with pytest.raises(ValueError, match="unknown"): + coerce_request(SimpleNamespace(**vars(captured[1]), unexpected_option=True)) + assert captured[0].family == "gemma" + assert captured[0].task == "text_generation" + assert captured[0].precision == ("fp16" if options else "fp32") + assert not output.exists() + + +def test_declared_help_is_offline_and_dependency_free(): + """Actual child-process help needs neither a checkpoint nor GPU imports.""" + import subprocess + import sys + + code = """ +import sys +from tensorrt_model_connect.family_cli import main +try: + main(["gemma", "build", "--help"]) +except SystemExit as error: + assert error.code == 0 +else: + raise AssertionError("help did not exit") +assert "families.gemma.cli" not in sys.modules +assert "tensorrt" not in sys.modules +assert "huggingface_hub" not in sys.modules +""" + result = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True, check=True) + assert "trtmc gemma build" in result.stdout diff --git a/website/docs/features/model-families.md b/website/docs/features/model-families.md index b9173b51fd..5f728b1339 100644 --- a/website/docs/features/model-families.md +++ b/website/docs/features/model-families.md @@ -68,6 +68,11 @@ wall-clock speedup claim. ### Gemma4 paired ONNX execution +Use `trtmc gemma build MODEL -o model.bundle` with the owning +family\u0027s options. `trtmc gemma build --help` works offline without +a checkpoint or GPU imports. This uses the existing +[family CLI protocol](../extend/family-cli.md), not an extension to the shared parser. + The Gemma family also owns explicit Gemma4-12B target/assistant MTP and Gemma4-12B/DSpark block7 execution through the optional pinned native Edge-LLM SDK. These are text-only FP16 paired profiles, qualified on SM80; selecting a From 2f4c0dc777fcd3543dee2e2ec119bb34cd4b8443 Mon Sep 17 00:00:00 2001 From: Joshua Calafato Date: Wed, 23 Sep 2026 21:21:03 +0000 Subject: [PATCH 5/6] fix(gemma): address Edge review findings Signed-off-by: Joshua Calafato --- families/gemma/edge_llm/README.md | 9 +++++++-- families/gemma/tests/test_model_type_gate.py | 6 ++++-- website/docs/features/model-families.md | 2 +- 3 files changed, 12 insertions(+), 5 deletions(-) diff --git a/families/gemma/edge_llm/README.md b/families/gemma/edge_llm/README.md index 87c85468b5..c7b66f8ed5 100644 --- a/families/gemma/edge_llm/README.md +++ b/families/gemma/edge_llm/README.md @@ -97,7 +97,8 @@ sampling equivalence or other model/platform combinations. The official Edge 0.10.1 static Gemma template omits the checkpoint's closed thought channel when thinking is disabled, and differs in enabled-thinking -system-prefix handling. A first uncorrected MTP run produced `thought\nParis` +system-prefix handling. A first uncorrected MTP run produced `thought +Paris` instead of the independent reference `Paris`, failing NED 0.6154 against 0.15. The family now validates the source single-user template during build and renders it faithfully before calling Edge with raw text. Unicode whitespace @@ -123,6 +124,10 @@ approved artifact cleanup, so replay requires rebuilding those exact profiles. ## Declared build command -This family uses the existing cli.json protocol introduced in #1310. The family\nowns its declaration, typed inputs and Python handler. The handler adapts those\ninputs to the unchanged builder API, preserving native/Edge dispatch and bundle\npublication. The legacy flat build command remains available for its existing\nordinary options; new family options use `trtmc gemma build`. +This family uses the existing cli.json protocol introduced in #1310. The family +owns its declaration, typed inputs and Python handler. The handler adapts those +inputs to the unchanged builder API, preserving native/Edge dispatch and bundle +publication. The legacy flat build command remains available for its existing +ordinary options; new family options use `trtmc gemma build`. Help is offline and does not need a local checkpoint. No shared parser hook or family registry entry is added. diff --git a/families/gemma/tests/test_model_type_gate.py b/families/gemma/tests/test_model_type_gate.py index ae1f950266..1a8b098d69 100644 --- a/families/gemma/tests/test_model_type_gate.py +++ b/families/gemma/tests/test_model_type_gate.py @@ -242,10 +242,12 @@ def paired(request, writer, execution): def test_bad_edge_cli_inputs_fail_before_backend_and_bundle(tmp_path, monkeypatch, options): from tensorrt_model_connect import family_cli as build_cli + from families.gemma import cli as owner + source = _model_dir(tmp_path / "target", "gemma4_unified") output = tmp_path / "pair.bundle" - monkeypatch.setattr(build_core, "_select_backend", lambda *_: pytest.fail("backend touched")) - monkeypatch.setattr(build_core, "BundleWriter", lambda *_: pytest.fail("writer created")) + monkeypatch.setattr(owner, "select_backend", lambda *_: pytest.fail("backend touched")) + monkeypatch.setattr(owner, "BundleWriter", lambda *_: pytest.fail("writer created")) with pytest.raises(ValueError): build_cli.main(["gemma", "build", str(source), "-o", str(output), *options]) assert not output.exists() diff --git a/website/docs/features/model-families.md b/website/docs/features/model-families.md index 5f728b1339..f94173e67e 100644 --- a/website/docs/features/model-families.md +++ b/website/docs/features/model-families.md @@ -69,7 +69,7 @@ wall-clock speedup claim. ### Gemma4 paired ONNX execution Use `trtmc gemma build MODEL -o model.bundle` with the owning -family\u0027s options. `trtmc gemma build --help` works offline without +family's options. `trtmc gemma build --help` works offline without a checkpoint or GPU imports. This uses the existing [family CLI protocol](../extend/family-cli.md), not an extension to the shared parser. From bea87abc7c5d8ef67ddf044dfc2e3f67bd9aad4c Mon Sep 17 00:00:00 2001 From: Joshua Calafato Date: Thu, 1 Oct 2026 19:44:18 +0000 Subject: [PATCH 6/6] fix(gemma): default paired CLI builds to fp16 Resolve omitted precision in the family handler: paired execution selects fp16, ordinary Gemma keeps fp32, and explicit values are never rewritten. Extend the existing routing test for omitted and explicit precision in both paired variants. Keep Edge admission and numerical gates unchanged. Signed-off-by: Joshua Calafato --- families/gemma/cli.json | 2 +- families/gemma/cli.py | 4 +++- families/gemma/tests/test_model_type_gate.py | 9 +++++++-- 3 files changed, 11 insertions(+), 4 deletions(-) diff --git a/families/gemma/cli.json b/families/gemma/cli.json index be801c5e83..80b55cf45e 100644 --- a/families/gemma/cli.json +++ b/families/gemma/cli.json @@ -50,7 +50,7 @@ "fp16", "bf16" ], - "default": "fp32" + "help": "Build precision (default: fp16 for paired execution, fp32 otherwise)" }, { "name": "backend", diff --git a/families/gemma/cli.py b/families/gemma/cli.py index efcd0fa2cb..50bc7912fb 100644 --- a/families/gemma/cli.py +++ b/families/gemma/cli.py @@ -36,13 +36,15 @@ def build_bundle(request: BuildRequest, output: Path) -> None: def build( *, model: str, output: Path, revision: str | None = None, - task: str = "text_generation", precision: str = "fp32", backend: str = "trt", + task: str = "text_generation", precision: str | None = None, backend: str = "trt", max_sequence_length: int | None = None, tensor_parallel_size: int = 1, verbose: bool = False, execution_variant: str | None = None, companion: list[str] | tuple[str, ...] = (), ) -> int: """Run the declared owner command; help never imports this handler.""" execution = execution_inputs(execution_variant, companion) + if precision is None: + precision = "fp16" if execution is not None else "fp32" model_dir = resolve_model(model, revision) resolve_family(load_model_metadata(model_dir), "gemma") request = BuildRequest( diff --git a/families/gemma/tests/test_model_type_gate.py b/families/gemma/tests/test_model_type_gate.py index 1a8b098d69..58065cd68e 100644 --- a/families/gemma/tests/test_model_type_gate.py +++ b/families/gemma/tests/test_model_type_gate.py @@ -205,7 +205,10 @@ def test_untyped_execution_fails_before_side_effects(tmp_path, monkeypatch): @pytest.mark.parametrize("variant", ["mtp", "dspark"]) -def test_edge_cli_routes_through_the_ordinary_family_entrypoint(tmp_path, monkeypatch, variant): +@pytest.mark.parametrize("precision", [None, "fp16", "fp32", "bf16"]) +def test_edge_cli_routes_through_the_ordinary_family_entrypoint( + tmp_path, monkeypatch, variant, precision +): from families.gemma.edge_llm import builder as edge_builder from tensorrt_model_connect import family_cli as build_cli @@ -218,13 +221,15 @@ def test_edge_cli_routes_through_the_ordinary_family_entrypoint(tmp_path, monkey def paired(request, writer, execution): assert isinstance(request, GemmaBuildRequest) assert request.execution is execution + assert request.precision == ("fp16" if precision is None else precision) seen.append(execution) writer.set_header(family="gemma", task=request.task, backend=request.backend) writer.add_json("edge-test.json", {"variant": execution.variant}) monkeypatch.setattr(edge_builder, "build", paired) + precision_args = [] if precision is None else ["--precision", precision] assert build_cli.main(["gemma", - "build", str(source), "--precision", "fp16", + "build", str(source), *precision_args, "-o", str(output), "--execution-variant", variant, "--companion", f"draft={draft}", ]) == 0