diff --git a/families/llama/checkpoint_mapper.py b/families/llama/checkpoint_mapper.py index 7fb1fb8b58..17a74ff53b 100644 --- a/families/llama/checkpoint_mapper.py +++ b/families/llama/checkpoint_mapper.py @@ -111,6 +111,8 @@ def load_standard_weights( f"Embedding shape {embedding.shape} != ({vocab}, {hidden})" ) weights["embedding"] = embedding.astype(target_dtype) + # Jiaxin Deng: keep only the stored embedding while loading decoder layers. + del embedding def _load_layer(layer_idx: int) -> tuple[int, WeightDict, int, int]: prefix = f"layer.{layer_idx}" @@ -128,36 +130,23 @@ def _load_layer(layer_idx: int) -> tuple[int, WeightDict, int, int]: layer[f"{prefix}.input_norm"] = input_norm.astype(np.float32) layer[f"{prefix}.post_attn_norm"] = post_norm.astype(np.float32) - # Q/K/V/O projections - q_raw = _load_tensor( - readers, _layer_key(layer_idx, "self_attn.q_proj.weight", model_prefix) - ) - k_raw = _load_tensor( - readers, _layer_key(layer_idx, "self_attn.k_proj.weight", model_prefix) - ) - v_raw = _load_tensor( - readers, _layer_key(layer_idx, "self_attn.v_proj.weight", model_prefix) - ) - o_raw = _load_tensor( - readers, _layer_key(layer_idx, "self_attn.o_proj.weight", model_prefix) - ) - - q_hidden = q_raw.shape[0] - gate_raw = _load_tensor( - readers, _layer_key(layer_idx, "mlp.gate_proj.weight", model_prefix) - ) - layer_mlp_size = gate_raw.shape[0] - - # Transpose all projections [out, in] -> [in, out] - q_t = _transpose_2d(q_raw, "q_proj", precision=layer_precision) - k_t = _transpose_2d(k_raw, "k_proj", precision=layer_precision) - v_t = _transpose_2d(v_raw, "v_proj", precision=layer_precision) - o_t = _transpose_2d(o_raw, "o_proj", precision=layer_precision) - - layer[f"{prefix}.w_q"] = q_t - layer[f"{prefix}.w_k"] = k_t - layer[f"{prefix}.w_v"] = v_t - layer[f"{prefix}.w_o"] = o_t + # Jiaxin Deng: release each FP32 source before loading the next projection. + for source, target in ( + ("self_attn.q_proj", "w_q"), + ("self_attn.k_proj", "w_k"), + ("self_attn.v_proj", "w_v"), + ("self_attn.o_proj", "w_o"), + ("mlp.gate_proj", "w_gate"), + ("mlp.up_proj", "w_up"), + ("mlp.down_proj", "w_down"), + ): + raw = _load_tensor(readers, _layer_key(layer_idx, f"{source}.weight", model_prefix)) + layer[f"{prefix}.{target}"] = _transpose_2d(raw, source, precision=layer_precision) + if target == "w_q": + q_hidden = raw.shape[0] + elif target == "w_gate": + layer_mlp_size = raw.shape[0] + del raw # Optional QKV biases (Qwen2 style) q_bias_key = _layer_key(layer_idx, "self_attn.q_proj.bias", model_prefix) @@ -182,16 +171,6 @@ def _load_layer(layer_idx: int) -> tuple[int, WeightDict, int, int]: _load_tensor(readers, k_norm_key).astype(np.float32), num_kv_heads ) - # MLP projections - up_raw = _load_tensor(readers, _layer_key(layer_idx, "mlp.up_proj.weight", model_prefix)) - down_raw = _load_tensor( - readers, _layer_key(layer_idx, "mlp.down_proj.weight", model_prefix) - ) - - layer[f"{prefix}.w_gate"] = _transpose_2d(gate_raw, "gate_proj", precision=layer_precision) - layer[f"{prefix}.w_up"] = _transpose_2d(up_raw, "up_proj", precision=layer_precision) - layer[f"{prefix}.w_down"] = _transpose_2d(down_raw, "down_proj", precision=layer_precision) - return layer_idx, layer, q_hidden, layer_mlp_size layer_results: list[tuple[int, WeightDict, int, int] | None] = [None] * num_layers @@ -232,7 +211,9 @@ def _load_layer(layer_idx: int) -> tuple[int, WeightDict, int, int]: ) else: # Tied embeddings - weights["w_out"] = _transpose_2d(embedding.copy(), "embedding_tied", precision=precision) + weights["w_out"] = _transpose_2d( + weights["embedding"].copy(), "embedding_tied", precision=precision + ) weights["_attention_size"] = attention_size # type: ignore[assignment] weights["_kv_attention_size"] = kv_attention_size # type: ignore[assignment] diff --git a/families/llama/model.py b/families/llama/model.py index 86ded25861..a366627a2b 100644 --- a/families/llama/model.py +++ b/families/llama/model.py @@ -261,6 +261,9 @@ def _build_native(request: "BuildRequest", writer: "BundleWriter") -> None: precision=precision, verbose=bool(request.verbose), ) + # Note (Jiaxin Deng): avoid retaining both plans during decode compilation. + writer.add_bytes("prefill.plan", prefill) + del prefill config.raw["_decoder_engine_role"] = "decode" decode = _build_engine( config, @@ -270,7 +273,6 @@ def _build_native(request: "BuildRequest", writer: "BundleWriter") -> None: verbose=bool(request.verbose), ) writer.add_bytes("engine.plan", decode) - writer.add_bytes("prefill.plan", prefill) layout = "split" config.raw.pop("_decoder_engine_role", None) diff --git a/families/llama/tests/test_checkpoint_cache.py b/families/llama/tests/test_checkpoint_cache.py new file mode 100644 index 0000000000..9f765a0ca5 --- /dev/null +++ b/families/llama/tests/test_checkpoint_cache.py @@ -0,0 +1,44 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Exercise staged checkpoint lookup without Hub metadata or network access.""" + +import pytest +from huggingface_hub import constants +from huggingface_hub.errors import LocalEntryNotFoundError + +from .test_e2e import _checkpoint + + +@pytest.fixture +def offline_cache(tmp_path, monkeypatch): + monkeypatch.setattr(constants, "HF_HUB_CACHE", str(tmp_path)) + monkeypatch.setattr(constants, "HF_HUB_OFFLINE", True) + return tmp_path + + +def test_pinned_checkpoint_uses_snapshot_without_cached_tree(offline_cache): + revision = "a" * 40 + snapshot = offline_cache / "models--trtmc-test--llama" / "snapshots" / revision + snapshot.mkdir(parents=True) + (snapshot / "config.json").write_text("{}", encoding="utf-8") + + assert _checkpoint({"hf_id": "trtmc-test/llama", "hf_revision": revision}) == snapshot + + +def test_missing_checkpoint_is_not_replaced_by_another_revision(offline_cache): + cached = offline_cache / "models--trtmc-test--llama" / "snapshots" / ("a" * 40) + cached.mkdir(parents=True) + (cached / "config.json").write_text("{}", encoding="utf-8") + + with pytest.raises(LocalEntryNotFoundError): + _checkpoint({"hf_id": "trtmc-test/llama", "hf_revision": "b" * 40}) + + +def test_staged_snapshot_still_requires_config(offline_cache): + revision = "a" * 40 + snapshot = offline_cache / "models--trtmc-test--llama" / "snapshots" / revision + snapshot.mkdir(parents=True) + + with pytest.raises(AssertionError): + _checkpoint({"hf_id": "trtmc-test/llama", "hf_revision": revision}) diff --git a/families/llama/tests/test_support.py b/families/llama/tests/test_support.py new file mode 100644 index 0000000000..762eb964b7 --- /dev/null +++ b/families/llama/tests/test_support.py @@ -0,0 +1,177 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""CPU coverage of build memory lifetime and bundle publication.""" + +import gc +import importlib.util +import json +from pathlib import Path +import struct +import sys +import types +import weakref + +import numpy as np +import pytest +from safetensors.numpy import save_file + +from tensorrt_model_connect import BuildRequest +from tensorrt_model_connect.bundle_writer import BundleWriter + +from .. import checkpoint_mapper +from ..config import ModelConfig + +@pytest.fixture +def request_and_writer(tmp_path, monkeypatch): + # Jiaxin Deng: stub only GPU compilation; exercise the real build and writer on CPU. + with monkeypatch.context() as imports: + for module_name, function_name in ( + ("dual_profile_decoder_builder", "build_dual_profile_decoder_engine"), + ("standard_decoder_builder", "build_standard_decoder_engine"), + ): + module = types.ModuleType(f"families.llama.{module_name}") + setattr(module, function_name, None) + imports.setitem(sys.modules, module.__name__, module) + spec = importlib.util.spec_from_file_location( + "families.llama._plan_lifetime_model", Path(__file__).parents[1] / "model.py" + ) + model = importlib.util.module_from_spec(spec) + spec.loader.exec_module(model) + checkpoint = tmp_path / "checkpoint" + checkpoint.mkdir() + (checkpoint / "config.json").write_text(json.dumps({ + "model_type": "llama", "hidden_size": 8, "intermediate_size": 16, + "num_hidden_layers": 1, "num_attention_heads": 2, + "num_key_value_heads": 1, "vocab_size": 32, + "max_position_embeddings": 32, "eos_token_id": 2, + })) + monkeypatch.setattr(model, "load_standard_weights", lambda *args, **kwargs: {}) + destination = tmp_path / "model.bundle" + request = BuildRequest( + model_dir=checkpoint, output_path=destination, family="llama", + task="text_generation", precision="fp16", max_sequence_length=16, + ) + writer = BundleWriter(destination) + yield model, request, writer + writer.abort() + + +def test_prefill_plan_is_released_before_decode_build(request_and_writer, monkeypatch): + model, request, writer = request_and_writer + released = [] + + class Plan(bytes): + def __del__(self): + released.append(True) + + def build_engine(config, *args, **kwargs): + if config.raw["_decoder_engine_role"] == "prefill": + return Plan(b"prefill bytes") + gc.collect() + assert released, "prefill plan is still resident during decode build" + return b"decode bytes" + + monkeypatch.setattr(model, "_build_engine", build_engine) + model.build(request, writer) + writer.finish() + data = request.output_path.read_bytes() + header_size = struct.unpack("