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..80b55cf45e --- /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" + ], + "help": "Build precision (default: fp16 for paired execution, fp32 otherwise)" + }, + { + "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..50bc7912fb --- /dev/null +++ b/families/gemma/cli.py @@ -0,0 +1,59 @@ +# 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 | 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( + 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 new file mode 100644 index 0000000000..c7b66f8ed5 --- /dev/null +++ b/families/gemma/edge_llm/README.md @@ -0,0 +1,133 @@ +# 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/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. + +Use the existing family CLI protocol with family-owned options (no checkpoint edits): + +```sh +trtmc gemma build /path/to/target --precision fp16 \ + -o /path/to/pair.bundle --execution-variant mtp \ + --companion draft=/path/to/assistant +``` + +`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. +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 +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 +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 +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. + +## 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. + +## Declared build command + +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/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/builder.py b/families/gemma/edge_llm/builder.py new file mode 100644 index 0000000000..7fb014df07 --- /dev/null +++ b/families/gemma/edge_llm/builder.py @@ -0,0 +1,243 @@ +# 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 for this family-owned offload.""" + 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/edge_llm/cli.py b/families/gemma/edge_llm/cli.py new file mode 100644 index 0000000000..467d8188d8 --- /dev/null +++ b/families/gemma/edge_llm/cli.py @@ -0,0 +1,29 @@ +# 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.""" + +from pathlib import Path + +from .config import BuildExecutionInputs, NamedCheckpoint + + +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 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 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(execution_variant, tuple(checkpoints)) diff --git a/families/gemma/edge_llm/config.py b/families/gemma/edge_llm/config.py new file mode 100644 index 0000000000..98a35b21fb --- /dev/null +++ b/families/gemma/edge_llm/config.py @@ -0,0 +1,89 @@ +# 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 ..build_request import BuildRequest, coerce_request + + +_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 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 2ab719e1c2..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,18 @@ 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: + 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") diff --git a/families/gemma/runtime/CMakeLists.txt b/families/gemma/runtime/CMakeLists.txt index 5303181cd9..9e7234f2a1 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,5 @@ if(TRTMC_BUILD_TESTS) endforeach() set_tests_properties(test_gemma_pipeline PROPERTIES SKIP_RETURN_CODE 77) 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..8f962b25c5 --- /dev/null +++ b/families/gemma/runtime/edge_llm/Adapter.cmake @@ -0,0 +1,22 @@ +# 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) + 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) + 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..aa338e9ee0 100644 --- a/families/gemma/tests/test_e2e.py +++ b/families/gemma/tests/test_e2e.py @@ -133,23 +133,26 @@ 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", ())) - 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, - ) + 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 ce0bcfb354..58065cd68e 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( @@ -68,12 +77,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) @@ -132,3 +141,223 @@ 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"]) +@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 + + 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 + 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_args, + "-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 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(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() + + +@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): + from families.gemma.edge_llm import cli + + assert cli.execution_inputs(None) is None + + +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)) + + +@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 cb3a6e18f6..f94173e67e 100644 --- a/website/docs/features/model-families.md +++ b/website/docs/features/model-families.md @@ -66,6 +66,28 @@ 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 + +Use `trtmc gemma build MODEL -o model.bundle` with the owning +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. + +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/README.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: