Skip to content

Commit 72b2536

Browse files
authored
Qualcomm AI Engine Direct - [GenAI Pipeline] PR5: Model preparation & quantization strategy implementations (pytorch#21899)
## Summary This PR implements the __model preparation__ and __quantization__ strategy implementations, replacing the `NotImplementedError` stubs with real logic. Each strategy delegates to injectable adapter interfaces (from PR4) for testability. ### What's included #### Strategy implementations (2 files + 1 `__init__` fix): - `ExecuTorchModelPreparationStrategy`: 5-step flow - `load_model` → `load_tokenizer` (via `ModelLoaderAdapter`) - `generate_calibration_data` (via separately-injectable `CalibrationDataAdapter`) - Optional tokenizer export for on-device runtime - Chat template extraction from tokenizer (with `extra_options` fallback) - Validates input config (`model_name`, `soc_model` required) - `ExecuTorchQuantizationStrategy`: Full PT2E single-graph pipeline via `QuantizerAdapter` - export → make_quantizer → prepare_pt2e → calibrate → convert_pt2e - Supports `quant_dtype`, `quant_recipe`, and per-channel options via `extra_options` - Handles any `Iterable` as calibration data (lists, DataLoaders, generators) - Validates calibration data is non-empty before export - Warns (does not fail) when `training_data` is provided (QAT deferred) - `strategies/model_preparation/__init__.py`: adds missing `ExecuTorchModelPreparationStrategy` import to `__all__` #### Unit tests: - `test_executorch_model_preparation_strategy.py` - `test_executorch_quantization_strategy.py` - `test_default_model_preparation_adapter.py` - `test_default_model_preparation_adapter.py` ### PR Review Checklist - All new classes follow single responsibility (one class per file) - Yes. - All dependencies are injected via constructor with sensible defaults - Yes. - All external calls are behind injectable interfaces - Yes (adapter pattern). - Unit tests cover every public method - Yes (100% coverage on strategy impls). - Docstrings on all public classes and methods - Yes. - Type annotations on all function signatures - Yes. - Logging follows the strategy in the LLD - Yes (info on entry/exit, debug per step). ### Related PRs - PR 1: Core data model, engine routing & exceptions: pytorch#20409 - PR 2: Strategy interfaces & stage wrappers: pytorch#20795 - PR 3: Pipeline orchestrator: pytorch#21149 - PR 4: Adapter interfaces, default implementations & dataset providers: pytorch#21751 - PR 5: Model preparation & quantization strategies: this pr. - PR 6: Compilation & inference strategy implementations: pending. - PR 7: Integration & E2E tests: pending. ## Test plan ### Run only tests added in this PR: ``` python -m pytest \ backends/qualcomm/genai_pipeline/tests/strategies/model_preparation/ \ backends/qualcomm/genai_pipeline/tests/strategies/quantization/ \ -v ``` ### Run only this PR's tests with coverage: ``` python -m pytest \ backends/qualcomm/genai_pipeline/tests/strategies/model_preparation/ \ backends/qualcomm/genai_pipeline/tests/strategies/quantization/ \ --cov=backends/qualcomm/genai_pipeline/strategies/model_preparation \ --cov=backends/qualcomm/genai_pipeline/strategies/quantization \ --cov-config=backends/qualcomm/.coveragerc \ --cov-report=term-missing ``` Result: ``` Name Stmts Miss Branch BrPart Cover Missing ---------------------------------------------------------------------------------------------------------------------------------------------------- backends/qualcomm/genai_pipeline/strategies/model_preparation/executorch_model_preparation_strategy.py 68 0 14 0 100% backends/qualcomm/genai_pipeline/strategies/model_preparation/model_loader_adapter.py 9 0 0 0 100% backends/qualcomm/genai_pipeline/strategies/model_preparation/model_preparation_strategy.py 7 0 0 0 100% backends/qualcomm/genai_pipeline/strategies/quantization/executorch_quantization_strategy.py 61 0 18 0 100% backends/qualcomm/genai_pipeline/strategies/quantization/quantization_strategy.py 7 0 0 0 100% backends/qualcomm/genai_pipeline/strategies/quantization/quantizer_adapter.py 9 0 0 0 100% ---------------------------------------------------------------------------------------------------------------------------------------------------- TOTAL 161 0 32 0 100% ``` ### Run all `genai_pipeline` tests: ``` python -m pytest backends/qualcomm/genai_pipeline/tests/ -v ``` ### Run all `genai_pipeline` tests with coverage: ``` python -m pytest backends/qualcomm/genai_pipeline/tests/ \ --cov=backends/qualcomm/genai_pipeline \ --cov-config=backends/qualcomm/.coveragerc \ --cov-report=term-missing ``` Result: ``` Name Stmts Miss Branch BrPart Cover Missing ---------------------------------------------------------------------------------------------------------------------------------------------------- backends/qualcomm/genai_pipeline/configs/compilation_input_config.py 11 0 0 0 100% backends/qualcomm/genai_pipeline/configs/compilation_output_config.py 8 0 0 0 100% backends/qualcomm/genai_pipeline/configs/inference_input_config.py 12 0 0 0 100% backends/qualcomm/genai_pipeline/configs/inference_output_config.py 9 0 0 0 100% backends/qualcomm/genai_pipeline/configs/model_preparation_input_config.py 7 0 0 0 100% backends/qualcomm/genai_pipeline/configs/model_preparation_output_config.py 12 0 0 0 100% backends/qualcomm/genai_pipeline/configs/quantization_input_config.py 13 0 0 0 100% backends/qualcomm/genai_pipeline/configs/quantization_output_config.py 5 0 0 0 100% backends/qualcomm/genai_pipeline/datasets/calibration_data_adapter.py 5 0 0 0 100% backends/qualcomm/genai_pipeline/datasets/default_calibration_data_adapter.py 26 0 6 0 100% backends/qualcomm/genai_pipeline/datasets/default_training_data_adapter.py 14 0 2 0 100% backends/qualcomm/genai_pipeline/datasets/training_data_adapter.py 5 0 0 0 100% backends/qualcomm/genai_pipeline/engine_proxy.py 20 0 4 0 100% backends/qualcomm/genai_pipeline/exceptions.py 20 0 6 0 100% backends/qualcomm/genai_pipeline/genai_pipeline.py 99 7 12 1 93% 190-201 backends/qualcomm/genai_pipeline/pipeline_context.py 52 0 14 0 100% backends/qualcomm/genai_pipeline/pipeline_stage.py 5 0 0 0 100% backends/qualcomm/genai_pipeline/stages/compilation_stage.py 14 0 0 0 100% backends/qualcomm/genai_pipeline/stages/inference_stage.py 14 0 0 0 100% backends/qualcomm/genai_pipeline/stages/model_preparation_stage.py 14 2 0 0 86% 30, 37 backends/qualcomm/genai_pipeline/stages/quantization_stage.py 14 0 0 0 100% backends/qualcomm/genai_pipeline/strategies/compilation/compilation_strategy.py 7 0 0 0 100% backends/qualcomm/genai_pipeline/strategies/compilation/compiler_adapter.py 11 0 0 0 100% backends/qualcomm/genai_pipeline/strategies/compilation/executorch_compilation_strategy.py 7 0 0 0 100% backends/qualcomm/genai_pipeline/strategies/inference/device_runner_adapter.py 14 0 0 0 100% backends/qualcomm/genai_pipeline/strategies/inference/executorch_inference_strategy.py 7 0 0 0 100% backends/qualcomm/genai_pipeline/strategies/inference/inference_strategy.py 7 0 0 0 100% backends/qualcomm/genai_pipeline/strategies/model_preparation/executorch_model_preparation_strategy.py 68 0 14 0 100% backends/qualcomm/genai_pipeline/strategies/model_preparation/model_loader_adapter.py 9 0 0 0 100% backends/qualcomm/genai_pipeline/strategies/model_preparation/model_preparation_strategy.py 7 0 0 0 100% backends/qualcomm/genai_pipeline/strategies/quantization/executorch_quantization_strategy.py 61 0 18 0 100% backends/qualcomm/genai_pipeline/strategies/quantization/quantization_strategy.py 7 0 0 0 100% backends/qualcomm/genai_pipeline/strategies/quantization/quantizer_adapter.py 9 0 0 0 100% ---------------------------------------------------------------------------------------------------------------------------------------------------- TOTAL 593 9 76 1 99% ```
1 parent 76434b8 commit 72b2536

15 files changed

Lines changed: 1367 additions & 52 deletions

‎backends/qualcomm/genai_pipeline/configs/model_preparation_output_config.py‎

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88

99
from dataclasses import dataclass
1010
from pathlib import Path
11-
from typing import Any, Iterable, Optional, TYPE_CHECKING
11+
from typing import Any, Iterable, Optional, Tuple, TYPE_CHECKING
1212

1313
if TYPE_CHECKING:
1414
from torch import nn
@@ -27,6 +27,12 @@ class ModelPreparationOutputConfig:
2727
Attributes:
2828
model_module: The prepared nn.Module ready for quantization.
2929
tokenizer: The tokenizer instance for encoding/decoding text.
30+
example_inputs: Positional example inputs for ``torch.export``, derived
31+
from the **model** (never from ``calibration_data``): they carry the
32+
exported graph's signature, its zero-initialized KV caches, and the
33+
AR length baked in because HTP has no dynamic shapes. The dependency
34+
runs model -> dataset, not the reverse -- the calibration dataset's
35+
attention-mask schema is itself derived from this tuple.
3036
calibration_data: Calibration samples. Any Iterable[Tuple[Tensor, ...]],
3137
including a DataLoader with a custom collate_fn.
3238
runtime_tokenizer_path: Path to the runtime tokenizer **file** (not the
@@ -36,6 +42,7 @@ class ModelPreparationOutputConfig:
3642

3743
model_module: Optional["nn.Module"] = None
3844
tokenizer: Any = None
45+
example_inputs: Optional[Tuple[Any, ...]] = None
3946
calibration_data: Optional[Iterable[Any]] = None
4047
runtime_tokenizer_path: Optional[Path] = None
4148
chat_template: Optional[str] = None

‎backends/qualcomm/genai_pipeline/configs/quantization_input_config.py‎

Lines changed: 17 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
from __future__ import annotations
88

99
from dataclasses import dataclass, field
10-
from typing import Any, Dict, Iterable, Optional, TYPE_CHECKING
10+
from typing import Any, Dict, Iterable, Optional, Tuple, TYPE_CHECKING
1111

1212
if TYPE_CHECKING:
1313
from executorch.backends.qualcomm.serialization.qc_schema import (
@@ -21,10 +21,11 @@
2121
class QuantizationInputConfig:
2222
"""Input configuration for the quantization stage.
2323
24-
``model_module`` and ``calibration_data`` are ``Optional`` only because the
25-
orchestrator builds this from the previous stage's output, which is empty
26-
when model preparation is skipped. **Both are required once the quantization
27-
stage executes**, and strategies should validate their presence.
24+
``model_module``, ``example_inputs`` and ``calibration_data`` are
25+
``Optional`` only because the orchestrator builds this from the previous
26+
stage's output, which is empty when model preparation is skipped. **All
27+
three are required once the quantization stage executes**, and strategies
28+
should validate their presence.
2829
2930
Flows needing no quantization (FP16, GPU backends) skip the stage entirely
3031
via ``GenAIPipeline.from_proxy(proxy, skip_stages={STAGE_QUANTIZATION})``
@@ -36,8 +37,17 @@ class QuantizationInputConfig:
3637
soc_model: The target SoC (e.g., QcomChipset.SM8750). Required.
3738
backend_type: QNN backend type (HTP, GPU, LPAI, etc.). Required.
3839
model_module: The nn.Module to quantize. Required when the stage runs.
40+
example_inputs: Positional example inputs for ``torch.export``. Required
41+
when the stage runs. Sourced from the **model** via
42+
``ModelLoaderAdapter.get_example_inputs``, never from
43+
``calibration_data``: this tuple defines the exported graph's
44+
positional signature, supplies the zero-initialized KV caches a
45+
dataset sample does not carry, and fixes the AR length because HTP
46+
has no dynamic shapes.
3947
calibration_data: Calibration samples. Required when the stage runs. Any
40-
Iterable[Tuple[Tensor, ...]], including a DataLoader.
48+
Iterable[Tuple[Tensor, ...]], including a DataLoader. Consumed only
49+
by ``calibrate()`` -- it is never indexed or peeked at, so a
50+
single-use generator stays intact.
4151
training_data: Training dataset for quantization-aware training (QAT),
4252
typically (features, labels) pairs. Mirrors ``qat_training_data`` in
4353
``build_executorch_binary``. ``None`` selects PTQ.
@@ -48,6 +58,7 @@ class QuantizationInputConfig:
4858
soc_model: "QcomChipset"
4959
backend_type: "QnnExecuTorchBackendType"
5060
model_module: Optional["nn.Module"] = None
61+
example_inputs: Optional[Tuple[Any, ...]] = None
5162
calibration_data: Optional[Iterable[Any]] = None
5263
training_data: Optional[Iterable[Any]] = None
5364
quant_recipe: Any = None

‎backends/qualcomm/genai_pipeline/genai_pipeline.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -213,6 +213,8 @@ def _run_quantization(
213213
soc_model=context.soc_model,
214214
backend_type=self._engine_proxy.backend_type,
215215
model_module=model_prep_output.model_module,
216+
# Export inputs come from the model, not from calibration_data.
217+
example_inputs=model_prep_output.example_inputs,
216218
calibration_data=model_prep_output.calibration_data,
217219
)
218220
output = self._quantization_stage.invoke(context, input_config)

‎backends/qualcomm/genai_pipeline/strategies/inference/default_device_runner_adapter.py‎

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -77,13 +77,18 @@ def execute(
7777
) -> InferenceResult:
7878
"""Execute the model on device via ADB.
7979
80+
Note: This is a two-step protocol. ``execute()`` runs the model and
81+
returns performance metrics. Call ``pull_results()`` afterward to
82+
retrieve the actual output data files from the device.
83+
8084
Args:
8185
inference_options: Engine-specific options. Supported keys:
8286
- ``method_index``: Index of the method to execute (default 0).
8387
- ``iteration``: Number of inference iterations (default 1).
8488
8589
Returns:
86-
InferenceResult with output data and performance metrics.
90+
InferenceResult with performance metrics. ``output_data`` is None
91+
until ``pull_results()`` is called separately.
8792
"""
8893
if self._adb is None:
8994
raise RuntimeError(

‎backends/qualcomm/genai_pipeline/strategies/model_preparation/__init__.py‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,9 @@
44
# This source code is licensed under the BSD-style license found in the
55
# LICENSE file in the root directory of this source tree.
66

7+
from executorch.backends.qualcomm.genai_pipeline.strategies.model_preparation.executorch_model_preparation_strategy import (
8+
ExecuTorchModelPreparationStrategy,
9+
)
710
from executorch.backends.qualcomm.genai_pipeline.strategies.model_preparation.model_loader_adapter import (
811
ModelLoaderAdapter,
912
)
@@ -12,6 +15,7 @@
1215
)
1316

1417
__all__ = [
18+
"ExecuTorchModelPreparationStrategy",
1519
"ModelLoaderAdapter",
1620
"ModelPreparationStrategy",
1721
]

‎backends/qualcomm/genai_pipeline/strategies/model_preparation/default_model_loader_adapter.py‎

Lines changed: 96 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88

99
import logging
1010
from pathlib import Path
11-
from typing import Any, Dict, Optional
11+
from typing import Any, Dict, Optional, Tuple
1212

1313
logger = logging.getLogger(__name__)
1414

@@ -30,6 +30,17 @@ class DefaultModelLoaderAdapter:
3030
``CalibrationDataAdapter``.
3131
"""
3232

33+
#: Batch size and sequence length of the generated example inputs. HTP has no
34+
#: dynamic shapes, so these dimensions are baked into the exported graph.
35+
DEFAULT_BATCH_SIZE = 1
36+
DEFAULT_AR_LEN = 1
37+
38+
#: Preferred runtime tokenizer file names, in priority order.
39+
#: ``pytorch_tokenizers.get_tokenizer`` dispatches on the file extension
40+
#: (``.json`` -> ``HuggingFaceTokenizer``, otherwise Llama2c/Tiktoken), so the
41+
#: file we hand back selects the runtime tokenizer implementation.
42+
RUNTIME_TOKENIZER_NAMES = ("tokenizer.json", "tokenizer.model")
43+
3344
def load_model(
3445
self,
3546
model_name: str,
@@ -94,6 +105,51 @@ def load_tokenizer(
94105
logger.info("Tokenizer loaded successfully")
95106
return tokenizer
96107

108+
def get_example_inputs(
109+
self,
110+
model: Any,
111+
extra_options: Optional[Dict[str, Any]] = None,
112+
) -> Tuple[Any, ...]:
113+
"""Build example inputs for ``torch.export`` from the model itself.
114+
115+
Prefers the model's own ``get_example_inputs()`` when it exposes one, so
116+
models that already describe their export signature (the LLM wrappers
117+
build a flat ``(tokens, attn_mask, pos_ids, *k_caches, *v_caches)``
118+
tuple) stay authoritative. Otherwise a minimal ``(input_ids,)`` is
119+
synthesized, which is the correct signature for a plain HuggingFace
120+
causal LM without an external KV cache.
121+
122+
Args:
123+
model: The module returned by :meth:`load_model`.
124+
extra_options: Additional options. Supported keys:
125+
- ``batch_size``: Batch dimension (default:
126+
``DEFAULT_BATCH_SIZE``).
127+
- ``ar_len``: Sequence length / autoregressive window
128+
(default: ``DEFAULT_AR_LEN``).
129+
130+
Returns:
131+
A flat tuple positionally matching ``model.forward``.
132+
"""
133+
import torch
134+
135+
extra_options = extra_options or {}
136+
137+
model_provided = getattr(model, "get_example_inputs", None)
138+
if callable(model_provided):
139+
logger.info("Using example inputs provided by the model")
140+
return tuple(model_provided())
141+
142+
batch_size = extra_options.get("batch_size", self.DEFAULT_BATCH_SIZE)
143+
ar_len = extra_options.get("ar_len", self.DEFAULT_AR_LEN)
144+
145+
logger.info(
146+
"Synthesizing example inputs with batch_size=%d, ar_len=%d",
147+
batch_size,
148+
ar_len,
149+
)
150+
# int64 token ids: the embedding lookup indexes with them.
151+
return (torch.zeros((batch_size, ar_len), dtype=torch.int64),)
152+
97153
def export_tokenizer(
98154
self,
99155
tokenizer: Any,
@@ -102,11 +158,19 @@ def export_tokenizer(
102158
) -> Path:
103159
"""Export tokenizer to disk and return the runtime tokenizer file.
104160
105-
``save_pretrained`` writes several files and returns the tuple of paths
106-
it wrote, with the tokenizer file last. Both ``llm::load_tokenizer`` and
107-
``pytorch_tokenizers.get_tokenizer`` expect that **single file**, not the
108-
containing directory, so we return it -- mirroring the existing
109-
``TokenizerWrapper._from_hf`` flow.
161+
``save_pretrained`` writes several files and returns the tuple of paths it
162+
wrote. Both ``llm::load_tokenizer`` and ``pytorch_tokenizers.get_tokenizer``
163+
expect a **single file**, not the containing directory, so one artifact has
164+
to be singled out.
165+
166+
The file is chosen **by name** -- ``tokenizer.json`` first, then
167+
``tokenizer.model`` -- rather than by position in the returned tuple.
168+
``get_tokenizer`` dispatches on the extension, so picking the wrong
169+
artifact silently constructs the wrong tokenizer class instead of raising,
170+
and ``save_pretrained``'s ordering is an implementation detail that varies
171+
with the tokenizer (fast vs slow, whether ``added_tokens.json`` is
172+
written). ``artifacts[-1]`` remains a last-resort fallback for tokenizers
173+
that emit neither name, mirroring ``TokenizerWrapper._from_hf``.
110174
111175
Args:
112176
tokenizer: The tokenizer instance to export.
@@ -128,6 +192,31 @@ def export_tokenizer(
128192
f"save_pretrained() reported no tokenizer artifacts in {output_dir}."
129193
)
130194

131-
runtime_tokenizer_path = Path(artifacts[-1])
195+
runtime_tokenizer_path = self._select_runtime_tokenizer(artifacts)
132196
logger.info("Tokenizer exported to %s", runtime_tokenizer_path)
133197
return runtime_tokenizer_path
198+
199+
@classmethod
200+
def _select_runtime_tokenizer(cls, artifacts: Any) -> Path:
201+
"""Pick the runtime tokenizer file out of ``save_pretrained``'s artifacts.
202+
203+
Args:
204+
artifacts: The paths reported by ``save_pretrained``.
205+
206+
Returns:
207+
The first artifact matching :attr:`RUNTIME_TOKENIZER_NAMES`, falling
208+
back to the last artifact when none matches.
209+
"""
210+
paths = [Path(artifact) for artifact in artifacts]
211+
212+
for name in cls.RUNTIME_TOKENIZER_NAMES:
213+
for path in paths:
214+
if path.name == name:
215+
return path
216+
217+
logger.warning(
218+
"None of %s found among tokenizer artifacts; falling back to %s.",
219+
", ".join(cls.RUNTIME_TOKENIZER_NAMES),
220+
paths[-1],
221+
)
222+
return paths[-1]

0 commit comments

Comments
 (0)