diff --git a/tests/models/composite.py b/tests/models/composite.py index 2bb8985..c49e127 100644 --- a/tests/models/composite.py +++ b/tests/models/composite.py @@ -40,6 +40,30 @@ def sdpa_externalize_spec(): ) +def rope_externalize_spec(): + """ExternalizeSpec targeting the RoPE composite.""" + from coreai_torch import ExternalizeSpec # noqa: PLC0415 + from coreai_torch.composite_ops import RoPE # noqa: PLC0415 + + return ExternalizeSpec( + target_class=RoPE, + composite_op_name="rope", + composite_attrs=["scale", "base", "dims", "interleaved"], + ) + + +def gathermm_externalize_spec(): + """ExternalizeSpec targeting the GatherMM composite.""" + from coreai_torch import ExternalizeSpec # noqa: PLC0415 + from coreai_torch.composite_ops import GatherMM # noqa: PLC0415 + + return ExternalizeSpec( + target_class=GatherMM, + composite_op_name="gather_mm", + composite_attrs=["num_batch_axes"], + ) + + class MNISTCompositeRMSNormModel(nn.Module): """Tiny MNIST classifier with an embedded RMSNormImpl composite op. @@ -130,3 +154,37 @@ def forward(self, x: torch.Tensor) -> torch.Tensor: v = v.unsqueeze(1) attn = self.composite(q, k, v).squeeze(1) return self.out(attn) + + +class CompositeRoPEModel(nn.Module): + """proj -> RoPE(composite) -> output_proj, fp16.""" + + def __init__(self, dim: int = 32) -> None: + from coreai_torch.composite_ops import RoPE # noqa: PLC0415 + + super().__init__() + self.proj_in = nn.Linear(dim, dim, bias=False) + self.composite = RoPE(dims=dim) + self.proj_out = nn.Linear(dim, dim, bias=False) + + def forward(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: + h = self.proj_in(x) + rotated = self.composite(h, cos, sin) + return self.proj_out(rotated) + + +class CompositeGatherMMModel(nn.Module): + """proj -> GatherMM(composite) -> output_proj, fp16.""" + + def __init__(self, dim: int = 32) -> None: + from coreai_torch.composite_ops import GatherMM # noqa: PLC0415 + + super().__init__() + self.proj_in = nn.Linear(dim, dim, bias=False) + self.composite = GatherMM() + self.proj_out = nn.Linear(dim, dim, bias=False) + + def forward(self, lhs: torch.Tensor, rhs: torch.Tensor) -> torch.Tensor: + h = self.proj_in(lhs) + out = self.composite(h, rhs) + return self.proj_out(out) diff --git a/tests/quantization/test_composite_op_externalize.py b/tests/quantization/test_composite_op_externalize.py index 87c4583..e98a071 100644 --- a/tests/quantization/test_composite_op_externalize.py +++ b/tests/quantization/test_composite_op_externalize.py @@ -13,6 +13,8 @@ from __future__ import annotations +from collections.abc import Callable + import pytest import torch import torch.nn as nn @@ -33,9 +35,13 @@ make_quant_config, ) from tests.models.composite import ( + CompositeGatherMMModel, CompositeRMSNormOnlyModel, + CompositeRoPEModel, CompositeSDPAModel, + gathermm_externalize_spec, rmsnorm_externalize_spec, + rope_externalize_spec, sdpa_externalize_spec, ) from tests.test_utils.general import ( @@ -46,6 +52,25 @@ ) +def _single_tensor_sample() -> tuple[torch.Tensor, ...]: + return (torch.randn(2, 4, 32, dtype=torch.float16),) + + +def _rope_sample() -> tuple[torch.Tensor, ...]: + return ( + torch.randn(2, 4, 32, dtype=torch.float16), + torch.randn(4, 16, dtype=torch.float16), + torch.randn(4, 16, dtype=torch.float16), + ) + + +def _gather_mm_sample() -> tuple[torch.Tensor, ...]: + return ( + torch.randn(4, 32, dtype=torch.float16), + torch.randn(32, 32, dtype=torch.float16), + ) + + @pytest.mark.parametrize( "quantize_activations", [ @@ -53,16 +78,29 @@ pytest.param(True, id="w8a8"), ], ) +@pytest.mark.parametrize( + "model_cls, spec, sample_fn", + [ + pytest.param(CompositeSDPAModel, sdpa_externalize_spec(), _single_tensor_sample, id="sdpa"), + pytest.param(CompositeRoPEModel, rope_externalize_spec(), _rope_sample, id="rope"), + pytest.param( + CompositeGatherMMModel, gathermm_externalize_spec(), _gather_mm_sample, id="gather-mm" + ), + ], +) def test_composite_op_survives_prepare_and_finalize( quantize_activations: bool, + model_cls: type[nn.Module], + spec: ExternalizeSpec, + sample_fn: Callable[[], tuple[torch.Tensor, ...]], ) -> None: """The externalized composite must remain a single opaque call_function node end-to-end, under both w8 and w8a8. """ - model = CompositeSDPAModel().eval().half() - sample = torch.randn(2, 4, 32, dtype=torch.float16) + model = model_cls().eval().half() + sample = sample_fn() - _patch_model_for_externalization(model, [sdpa_externalize_spec()]) + _patch_model_for_externalization(model, [spec]) op_name = model.composite._externalize_op_name target_substr = f"coreai_torch_ext.{op_name}" @@ -74,7 +112,7 @@ def test_composite_op_survives_prepare_and_finalize( execution_mode="graph", ), ) - prepared = quantizer.prepare((sample,)) + prepared = quantizer.prepare(sample) assert_single_call_function_node(prepared, target_substr, stage="prepared") finalized = quantizer.finalize(backend=ExportBackend.CoreAI) @@ -121,7 +159,7 @@ def _config( def _quantize_with_externalization_and_verify( self, model: nn.Module, - sample: torch.Tensor, + sample: tuple[torch.Tensor, ...], spec: ExternalizeSpec, module_name: str, config: QuantizerConfig, @@ -139,7 +177,7 @@ def _quantize_with_externalization_and_verify( target_substr = f"coreai_torch_ext.{op_name}" quantizer = Quantizer(model, config) - prepared = quantizer.prepare((sample,)) + prepared = quantizer.prepare(sample) assert_single_call_function_node(prepared, target_substr, stage="prepared") finalized = quantizer.finalize(backend=ExportBackend.CoreAI) @@ -184,15 +222,15 @@ def _assert_boundary_quantized( @pytest.mark.parametrize("target_by", ["name", "type"]) @pytest.mark.parametrize( - # (model class, externalize spec, submodule attribute name, tensor input count). - # Both models default to dim=32 and accept the same rank-3 fp16 sample. - "model_cls, spec, module_name, num_tensor_inputs", + # (model class, externalize spec, submodule attribute name, tensor input count, sample_fn). + "model_cls, spec, module_name, num_tensor_inputs, sample_fn", [ pytest.param( CompositeRMSNormOnlyModel, rmsnorm_externalize_spec(), "norm", 1, + _single_tensor_sample, id="rmsnorm-only", ), pytest.param( @@ -200,8 +238,25 @@ def _assert_boundary_quantized( sdpa_externalize_spec(), "composite", 3, + _single_tensor_sample, id="sdpa-qkv", ), + pytest.param( + CompositeRoPEModel, + rope_externalize_spec(), + "composite", + 3, + _rope_sample, + id="rope", + ), + pytest.param( + CompositeGatherMMModel, + gathermm_externalize_spec(), + "composite", + 2, + _gather_mm_sample, + id="gather-mm", + ), ], ) def test_composite_boundary_quantized( @@ -210,12 +265,13 @@ def test_composite_boundary_quantized( spec: ExternalizeSpec, module_name: str, num_tensor_inputs: int, + sample_fn: Callable[[], tuple[torch.Tensor, ...]], target_by: str, ) -> None: # uint8 boundary edges stay distinguishable from the int8 global ones. boundary_dtype = torch.uint8 model = model_cls().eval().half() - sample = torch.randn(2, 4, 32, dtype=torch.float16) + sample = sample_fn() config = self._config(spec, module_name, target_by, boundary_dtype) finalized, target_substr = self._quantize_with_externalization_and_verify( model, sample, spec, module_name, config @@ -223,7 +279,44 @@ def test_composite_boundary_quantized( self._assert_boundary_quantized(finalized, target_substr, num_tensor_inputs, boundary_dtype) @pytest.mark.parametrize("target_by", ["name", "type"]) - def test_composite_boundary_input_index_selects_those_args(self, target_by: str) -> None: + @pytest.mark.parametrize( + "model_cls, spec, num_tensor_inputs, quantized_indices, sample_fn", + [ + pytest.param( + CompositeSDPAModel, + sdpa_externalize_spec(), + 3, + (0, 2), + _single_tensor_sample, + id="sdpa", + ), + pytest.param( + CompositeRoPEModel, + rope_externalize_spec(), + 3, + (0, 2), + _rope_sample, + id="rope", + ), + pytest.param( + CompositeGatherMMModel, + gathermm_externalize_spec(), + 2, + (0,), + _gather_mm_sample, + id="gather-mm", + ), + ], + ) + def test_composite_boundary_input_index_selects_those_args( + self, + target_by: str, + model_cls: type[nn.Module], + spec: ExternalizeSpec, + num_tensor_inputs: int, + quantized_indices: tuple[int, ...], + sample_fn: Callable[[], tuple[torch.Tensor, ...]], + ) -> None: """Integer keys in ``module_input_spec`` quantize exactly those positional args for composite ops. @@ -233,11 +326,8 @@ def test_composite_boundary_input_index_selects_those_args(self, target_by: str) """ # uint8 boundary edges stay distinguishable from the int8 global ones. boundary_dtype = torch.uint8 - quantized_indices = (0, 2) - num_tensor_inputs = 3 - model = CompositeSDPAModel().eval().half() - sample = torch.randn(2, 4, 32, dtype=torch.float16) - spec = sdpa_externalize_spec() + model = model_cls().eval().half() + sample = sample_fn() config = self._config( spec, "composite",