Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
58 changes: 58 additions & 0 deletions tests/models/composite.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand Down Expand Up @@ -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)
122 changes: 106 additions & 16 deletions tests/quantization/test_composite_op_externalize.py
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,8 @@

from __future__ import annotations

from collections.abc import Callable

import pytest
import torch
import torch.nn as nn
Expand All @@ -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 (
Expand All @@ -46,23 +52,55 @@
)


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",
[
pytest.param(False, id="w8-weight-only"),
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}"

Expand All @@ -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)
Expand Down Expand Up @@ -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,
Expand All @@ -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)
Expand Down Expand Up @@ -184,24 +222,41 @@ 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(
CompositeSDPAModel,
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(
Expand All @@ -210,20 +265,58 @@ 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
)
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.

Expand All @@ -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",
Expand Down