diff --git a/families/gemma/model.py b/families/gemma/model.py index bb6b252cd..2ab719e1c 100644 --- a/families/gemma/model.py +++ b/families/gemma/model.py @@ -359,6 +359,40 @@ def _runtime_config(model_dir: Path, config: ModelConfig, **updates) -> dict: return runtime +# The split layout stores the weights twice: a prefill engine with a dynamic +# sequence axis plus a decode engine with a static Sq=1 graph. That duplication +# buys decode throughput - measured with apps/benchmark at about 6% on +# gemma-3-4b over 200 tokens - and it is worth paying while the pair fits. +# +# It stops being a trade when the pair cannot be loaded. The qualification +# target is an L40S, which reports 46068 MiB, so about 45 GiB total and roughly +# 44 GiB once the driver and CUDA context are accounted for. +# +# 19 GiB per engine puts a pair at 38 GiB and leaves about 6 GiB for the KV +# cache, activations and TensorRT scratch. Measured with the estimator below: +# gemma-2-2b 4.9 GiB, gemma-3-270m 0.5 GiB, gemma-3-1b 2.0 GiB, gemma-3-4b +# 8.9 GiB, gemma-3-12b 25.0 GiB, gemma-3-27b 55.6 GiB. So everything up to 4b +# keeps the split pair, while 12b and 27b build a single plan - 12b's pair is +# 50 GiB and would not load on the target device at all. +# +# Erring low is safe and erring high is not: falling back unnecessarily costs +# about 6% of decode, while staying on split when the pair does not fit cannot +# be loaded. +_MAX_SPLIT_ENGINE_BYTES = 19 * 1024**3 + +# A serialized plan runs a little over the raw weights. Checked against two +# measured points: gemma-3-4b's engines are 8.5 GiB each and this returns +# 8.9 GiB; TensorRT asked 52.9 GiB for gemma-3-27b and this returns 55.6 GiB. +_PLAN_OVERHEAD = 1.05 + + +def _decoder_engine_bytes(weights: "WeightDict", precision: str) -> int: + """Roughly what one decoder engine will occupy on the device.""" + element = 2 if str(precision).lower() in {"fp16", "bf16"} else 4 + parameters = sum(int(getattr(value, "size", 0)) for value in weights.values()) + return int(parameters * element * _PLAN_OVERHEAD) + + def build(request: "BuildRequest", writer: "BundleWriter") -> None: """Build one Gemma bundle through family-owned code only.""" if request.dynamic_kv_cache: @@ -441,6 +475,21 @@ def build(request: "BuildRequest", writer: "BundleWriter") -> None: ) writer.add_bytes(f"engine.rank{rank}.plan", plan) layout = "dual_profile" + elif _decoder_engine_bytes(weights, precision) > _MAX_SPLIT_ENGINE_BYTES: + # One plan carrying both profiles, so the weights are stored once. + config.raw["_decoder_engine_role"] = "dual_profile" + plan = model.build_engine( + config, + weights, + max_sequence_length, + precision=precision, + quant_ctx=None, + verbose=bool(request.verbose), + parallel_config=parallel, + ) + config.raw.pop("_decoder_engine_role", None) + writer.add_bytes("engine.plan", plan) + layout = "dual_profile" else: config.raw["_decoder_engine_role"] = "prefill" prefill = model.build_engine( diff --git a/families/gemma/tests/manifests/gemma-3-27b.json b/families/gemma/tests/manifests/gemma-3-27b.json new file mode 100644 index 000000000..a3632aa97 --- /dev/null +++ b/families/gemma/tests/manifests/gemma-3-27b.json @@ -0,0 +1,22 @@ +{ + "name": "gemma-3-27b", + "hf_id": "google/gemma-3-27b-it", + "bundle": "gemma-3-27b.bundle", + "family": "gemma", + "task": "text_generation", + "precision": "bf16", + "trust_remote_code": false, + "testcases": [ + { + "name": "gemma-3-27b", + "premerge": true, + "reference_precision": "fp32", + "prompt": "What is the capital of France? Answer with just the name.", + "max_new_tokens": 8, + "use_chat_template": true, + "enable_thinking": false + } + ], + "max_sequence_length": 256, + "tensor_parallel_size": 1 +} diff --git a/families/gemma/tests/test_split_engine_budget.py b/families/gemma/tests/test_split_engine_budget.py new file mode 100644 index 000000000..cefa680da --- /dev/null +++ b/families/gemma/tests/test_split_engine_budget.py @@ -0,0 +1,53 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""A Gemma too large for a split pair must build one dual-profile plan instead. + +The split layout keeps a separate decode engine with a static Sq=1 graph, which +costs a second full copy of the weights on the device. The qualification target +is an L40S at about 45 GiB, so a pair is only affordable up to roughly 19 GiB +per engine: gemma-3-12b needs 25 GiB each and its 50 GiB pair does not fit. +""" + +from __future__ import annotations + +import numpy as np + +from families.gemma.model import ( + _MAX_SPLIT_ENGINE_BYTES, + _decoder_engine_bytes, +) + + +def _weights(parameters: int) -> dict: + return {"w": np.zeros(parameters, dtype=np.float32)} + + +def test_bytes_follow_the_build_precision(): + half = _decoder_engine_bytes(_weights(1_000_000), "bf16") + full = _decoder_engine_bytes(_weights(1_000_000), "fp32") + + assert full == 2 * half + # 1e6 parameters at 2 bytes, plus the measured plan overhead. + assert half == int(1_000_000 * 2 * 1.05) + + +def test_the_small_widths_stay_on_split(): + """Parameter counts of the shipped text decoders that still fit a pair.""" + for parameters in (0.27e9, 1.00e9, 2.61e9, 3.88e9): + assert _decoder_engine_bytes(_weights(int(parameters)), "bf16") <= _MAX_SPLIT_ENGINE_BYTES + + +def test_12b_and_27b_exceed_the_split_budget(): + for parameters in (11.77e9, 27.01e9): + assert _decoder_engine_bytes(_weights(int(parameters)), "bf16") > _MAX_SPLIT_ENGINE_BYTES + + +def test_a_pair_at_the_budget_fits_the_qualification_target(): + """Both engines plus working memory have to fit an L40S, not just the weights. + + The card reports 46068 MiB, so about 45 GiB; a pair at the budget is 38 GiB. + """ + l40s_bytes = 46068 * 1024**2 + assert 2 * _MAX_SPLIT_ENGINE_BYTES < l40s_bytes + assert l40s_bytes - 2 * _MAX_SPLIT_ENGINE_BYTES >= 5 * 1024**3 diff --git a/qualification_tests/benchmark_qualification/performance/config/release.yaml b/qualification_tests/benchmark_qualification/performance/config/release.yaml index 5ebbc0e3e..607b150bc 100644 --- a/qualification_tests/benchmark_qualification/performance/config/release.yaml +++ b/qualification_tests/benchmark_qualification/performance/config/release.yaml @@ -96,6 +96,8 @@ excluded_profiles: reason: *gemma3_performance_exclusion - model: gemma-3-12b reason: *gemma3_performance_exclusion + - model: gemma-3-27b + reason: *gemma3_performance_exclusion - model: gemma-3-1b reason: >- Functional and Hugging Face reference-parity qualification is present,