Skip to content

Commit a761ed0

Browse files
committed
test(axis-defaults): add graph-mode regression tests for ConvTranspose2d/3d axis resolution
1 parent bb9b589 commit a761ed0

3 files changed

Lines changed: 91 additions & 71 deletions

File tree

src/coreai_opt/_utils/torch_utils.py

Lines changed: 2 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -40,19 +40,12 @@
4040
torch.ops.aten.conv2d.default: torch.nn.Conv2d,
4141
torch.ops.aten.conv3d.default: torch.nn.Conv3d,
4242
torch.ops.aten.conv_transpose1d.default: torch.nn.ConvTranspose1d,
43-
torch.ops.aten.conv_transpose2d.default: torch.nn.ConvTranspose2d,
44-
torch.ops.aten.conv_transpose3d.default: torch.nn.ConvTranspose3d,
43+
torch.ops.aten.conv_transpose2d.input: torch.nn.ConvTranspose2d,
44+
torch.ops.aten.conv_transpose3d.input: torch.nn.ConvTranspose3d,
4545
torch.ops.aten.linear.default: torch.nn.Linear,
4646
torch.ops.aten.embedding.default: torch.nn.Embedding,
4747
}
4848

49-
# Backward-compat aliases for torch < 2.9
50-
try:
51-
ATEN_OP_TO_MODULE_TYPE[torch.ops.aten.conv_transpose2d.input] = torch.nn.ConvTranspose2d
52-
ATEN_OP_TO_MODULE_TYPE[torch.ops.aten.conv_transpose3d.input] = torch.nn.ConvTranspose3d
53-
except AttributeError:
54-
pass
55-
5649

5750
class NamedModule(NamedTuple):
5851
"""NamedTuple for holding name and module info"""

tests/quantization/test_axis_defaults.py

Lines changed: 89 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -422,3 +422,92 @@ def test_different_fqs_independent(self):
422422
_apply_defaults(fq_map)
423423
assert fq_linear.granularity.axis == 0
424424
assert fq_conv_t.granularity.axis == 0
425+
426+
427+
class TestConvTransposeAxisDefaultsGraph:
428+
"""Graph-mode regression tests for ConvTranspose2d/3d axis default resolution.
429+
430+
``ATEN_OP_TO_MODULE_TYPE`` maps the aten op emitted in the exported graph
431+
(e.g. ``aten.conv_transpose2d.input``) to the corresponding ``nn.Module``
432+
type so that ``_apply_defaults`` can look up the correct weight axis from
433+
``_WEIGHT_AXIS_SPECS``. These tests verify the full coreai-opt path:
434+
Quantizer.prepare → axis-defaults pass → correct axis on weight FQ.
435+
436+
ConvTranspose weight layout is ``[in_ch, out_ch, ...]``, so:
437+
- per-channel axis (output channels) = 1
438+
- per-block axis (input channels) = 0
439+
"""
440+
441+
@pytest.mark.parametrize(
442+
("make_model", "make_input"),
443+
[
444+
pytest.param(
445+
lambda: nn.ConvTranspose2d(16, 8, 3, padding=1),
446+
lambda: torch.randn(1, 16, 8, 8),
447+
id="conv_transpose2d",
448+
),
449+
pytest.param(
450+
lambda: nn.ConvTranspose3d(16, 8, 3, padding=1),
451+
lambda: torch.randn(1, 16, 4, 4, 4),
452+
id="conv_transpose3d",
453+
),
454+
],
455+
)
456+
@pytest.mark.parametrize(
457+
("granularity", "expected_axis"),
458+
[
459+
pytest.param(PerChannelGranularity(axis=None), 1, id="per_channel_axis_1"),
460+
pytest.param(
461+
PerBlockGranularity(axis=None, block_size=_TEST_BLOCK_SIZE), 0, id="per_block_axis_0"
462+
),
463+
],
464+
)
465+
def test_axis_none_resolves_for_conv_transpose(
466+
self,
467+
make_model,
468+
make_input,
469+
granularity,
470+
expected_axis,
471+
):
472+
"""ConvTranspose axis=None resolves to the correct default in graph mode.
473+
474+
Per-channel should resolve to axis 1 (output channels), per-block to
475+
axis 0 (input channels), reflecting the [in_ch, out_ch, ...] weight layout.
476+
"""
477+
config = _make_config(granularity, execution_mode="graph")
478+
prepared = Quantizer(make_model(), config).prepare((make_input(),))
479+
480+
weight_fqs = _get_weight_fqs(prepared)
481+
assert len(weight_fqs) == 1
482+
assert weight_fqs[0].granularity.axis == expected_axis
483+
484+
@pytest.mark.parametrize(
485+
("make_model", "make_input"),
486+
[
487+
pytest.param(
488+
lambda: nn.ConvTranspose2d(16, 8, 3, padding=1),
489+
lambda: torch.randn(1, 16, 8, 8),
490+
id="conv_transpose2d",
491+
),
492+
pytest.param(
493+
lambda: nn.ConvTranspose3d(16, 8, 3, padding=1),
494+
lambda: torch.randn(1, 16, 4, 4, 4),
495+
id="conv_transpose3d",
496+
),
497+
],
498+
)
499+
def test_prepare_calibrate_finalize_conv_transpose_graph(self, make_model, make_input):
500+
"""Full graph-mode workflow succeeds for ConvTranspose with axis=None."""
501+
config = _make_config(PerChannelGranularity(axis=None), execution_mode="graph")
502+
quantizer = Quantizer(make_model(), config)
503+
example_input = make_input()
504+
505+
prepared = quantizer.prepare((example_input,))
506+
with quantizer.calibration_mode():
507+
prepared(example_input)
508+
quantizer.finalize()
509+
510+
weight_fqs = _get_weight_fqs(prepared)
511+
assert len(weight_fqs) == 1
512+
assert weight_fqs[0].granularity.axis is not None
513+

tests/test_utils/test_torch_utils.py

Lines changed: 0 additions & 62 deletions
Original file line numberDiff line numberDiff line change
@@ -13,13 +13,11 @@
1313

1414
from coreai_opt._utils.fx_utils import normalize_module_fqn
1515
from coreai_opt._utils.torch_utils import (
16-
ATEN_OP_TO_MODULE_TYPE,
1716
mmap_module_state_dict,
1817
move_model_to_eval,
1918
move_model_to_train,
2019
normalize_axis,
2120
)
22-
from coreai_opt._utils.version_utils import version_ge
2321

2422

2523
class TestMoveModelContextManagers:
@@ -165,63 +163,3 @@ def test_raises_on_non_cpu_tensor(tmp_path):
165163

166164
with pytest.raises(ValueError, match="requires CPU tensors"):
167165
mmap_module_state_dict(model, tmp_path / "model.safetensors")
168-
169-
170-
class TestAtenOpToModuleType:
171-
"""Tests for ATEN_OP_TO_MODULE_TYPE overload correctness across torch versions.
172-
173-
Regression test for: conv_transpose2d/3d were mapped to the `.input` overload
174-
which was deprecated in PyTorch 2.9. The mapping must use `.default` as the
175-
primary key while retaining the `.input` alias for backward compat on < 2.9.
176-
"""
177-
178-
@staticmethod
179-
def test_conv_transpose2d_default_overload_present():
180-
"""conv_transpose2d.default must always be in the mapping (torch >= 2.9 primary path)."""
181-
assert torch.ops.aten.conv_transpose2d.default in ATEN_OP_TO_MODULE_TYPE
182-
assert ATEN_OP_TO_MODULE_TYPE[torch.ops.aten.conv_transpose2d.default] is nn.ConvTranspose2d
183-
184-
@staticmethod
185-
def test_conv_transpose3d_default_overload_present():
186-
"""conv_transpose3d.default must always be in the mapping (torch >= 2.9 primary path)."""
187-
assert torch.ops.aten.conv_transpose3d.default in ATEN_OP_TO_MODULE_TYPE
188-
assert ATEN_OP_TO_MODULE_TYPE[torch.ops.aten.conv_transpose3d.default] is nn.ConvTranspose3d
189-
190-
@staticmethod
191-
def test_conv_transpose2d_input_overload_compat():
192-
"""conv_transpose2d.input alias is present on torch < 2.9 for backward compat.
193-
194-
On torch >= 2.9 the `.input` overload may not exist at all; the test is
195-
skipped in that case since the `.default` path fully covers those versions.
196-
"""
197-
if not version_ge(torch, "2.9"):
198-
# On < 2.9 the .input overload exists and must also resolve to ConvTranspose2d
199-
assert torch.ops.aten.conv_transpose2d.input in ATEN_OP_TO_MODULE_TYPE
200-
assert ATEN_OP_TO_MODULE_TYPE[torch.ops.aten.conv_transpose2d.input] is nn.ConvTranspose2d
201-
else:
202-
# On >= 2.9 the .input overload may be absent; that's expected and fine
203-
pytest.skip("conv_transpose2d.input overload not expected on torch >= 2.9")
204-
205-
@staticmethod
206-
def test_conv_transpose3d_input_overload_compat():
207-
"""conv_transpose3d.input alias is present on torch < 2.9 for backward compat."""
208-
if not version_ge(torch, "2.9"):
209-
assert torch.ops.aten.conv_transpose3d.input in ATEN_OP_TO_MODULE_TYPE
210-
assert ATEN_OP_TO_MODULE_TYPE[torch.ops.aten.conv_transpose3d.input] is nn.ConvTranspose3d
211-
else:
212-
pytest.skip("conv_transpose3d.input overload not expected on torch >= 2.9")
213-
214-
@staticmethod
215-
def test_all_other_ops_unaffected():
216-
"""Sanity-check that unrelated ops in the mapping were not disturbed."""
217-
expected = {
218-
torch.ops.aten.conv1d.default: nn.Conv1d,
219-
torch.ops.aten.conv2d.default: nn.Conv2d,
220-
torch.ops.aten.conv3d.default: nn.Conv3d,
221-
torch.ops.aten.conv_transpose1d.default: nn.ConvTranspose1d,
222-
torch.ops.aten.linear.default: nn.Linear,
223-
torch.ops.aten.embedding.default: nn.Embedding,
224-
}
225-
for op, module_type in expected.items():
226-
assert op in ATEN_OP_TO_MODULE_TYPE
227-
assert ATEN_OP_TO_MODULE_TYPE[op] is module_type

0 commit comments

Comments
 (0)