|
4 | 4 | from collections.abc import Callable |
5 | 5 | from pathlib import Path |
6 | 6 | from typing import Any |
| 7 | +from unittest.mock import ANY, patch |
7 | 8 |
|
8 | 9 | import hypothesis.strategies as st |
9 | 10 | import numpy as np |
|
17 | 18 | from tket._ops import TketOp |
18 | 19 | from tket._pattern import Rule, RuleMatcher |
19 | 20 | from tket._state import CompilationState |
| 21 | +from tket._state.build import H, from_coms |
| 22 | +from tket._tket import passes as rust_passes |
20 | 23 | from tket.passes import ( |
21 | 24 | GlobalScope, |
22 | 25 | InlineFunctions, |
23 | 26 | ModifierResolverPass, |
24 | 27 | Normalize, |
25 | 28 | NormalizeGuppy, |
| 29 | + ParallelMode, |
| 30 | + PauliGraphResynthesis, |
26 | 31 | PlatformTarget, |
27 | 32 | PytketHugrPass, |
28 | 33 | QSystemRebasePass, |
@@ -507,3 +512,94 @@ def test_python_qsystem_pass_with_modifiers() -> None: |
507 | 512 | except Exception as exc: # noqa: BLE001 |
508 | 513 | failures.append(f"{hugr_path}: {exc}") |
509 | 514 | assert not failures, "QSystem pass failures:\n" + "\n".join(failures) |
| 515 | + |
| 516 | + |
| 517 | +def _count_hadamards(hugr: Hugr) -> int: |
| 518 | + return sum( |
| 519 | + data.op.name().rsplit(".", maxsplit=1)[-1] == "H" for _, data in hugr.nodes() |
| 520 | + ) |
| 521 | + |
| 522 | + |
| 523 | +def test_resynthesis_with_default_options() -> None: |
| 524 | + hugr = from_coms(H(0), H(0)).to_python().modules[0] |
| 525 | + optimisation = PauliGraphResynthesis() |
| 526 | + |
| 527 | + result = optimisation.run(hugr, inplace=False) |
| 528 | + |
| 529 | + assert optimisation.parallel_mode is ParallelMode.Auto |
| 530 | + assert optimisation.window_size is None |
| 531 | + assert optimisation.pool_size is None |
| 532 | + assert optimisation.top_up_size is None |
| 533 | + assert optimisation.seed is None |
| 534 | + assert result.results == [("PauliGraphResynthesis", None)] |
| 535 | + assert _count_hadamards(result.hugr) == 0 |
| 536 | + assert _count_hadamards(hugr) == 2 |
| 537 | + CompilationState.from_python(result.hugr).validate() |
| 538 | + |
| 539 | + |
| 540 | +def test_custom_options_and_scope() -> None: |
| 541 | + hugr = from_coms(H(0), H(0)).to_python().modules[0] |
| 542 | + optimisation = PauliGraphResynthesis( |
| 543 | + window_size=16, |
| 544 | + pool_size=32, |
| 545 | + top_up_size=4, |
| 546 | + seed=7, |
| 547 | + parallel_mode=ParallelMode.On, |
| 548 | + ) |
| 549 | + assert optimisation.with_scope(GlobalScope.PRESERVE_ALL) is optimisation |
| 550 | + |
| 551 | + with patch.object( |
| 552 | + rust_passes, |
| 553 | + "pauli_graph_resynthesis", |
| 554 | + wraps=rust_passes.pauli_graph_resynthesis, |
| 555 | + ) as resynthesis: |
| 556 | + result = optimisation.run(hugr, inplace=False) |
| 557 | + |
| 558 | + resynthesis.assert_called_once_with( |
| 559 | + ANY, |
| 560 | + scope=GlobalScope.PRESERVE_ALL, |
| 561 | + window_size=16, |
| 562 | + pool_size=32, |
| 563 | + top_up_size=4, |
| 564 | + seed=7, |
| 565 | + parallel_mode=ParallelMode.On, |
| 566 | + ) |
| 567 | + assert _count_hadamards(result.hugr) == 0 |
| 568 | + |
| 569 | + |
| 570 | +@pytest.mark.parametrize( |
| 571 | + ("parameter", "value", "message"), |
| 572 | + [ |
| 573 | + ("window_size", -1, "window_size must be positive"), |
| 574 | + ("window_size", 0, "window_size must be positive"), |
| 575 | + ("pool_size", -1, "pool_size must be positive"), |
| 576 | + ("pool_size", 0, "pool_size must be positive"), |
| 577 | + ("top_up_size", -1, "top_up_size must be positive"), |
| 578 | + ("top_up_size", 0, "top_up_size must be positive"), |
| 579 | + ("seed", -1, "seed must be non-negative"), |
| 580 | + ], |
| 581 | +) |
| 582 | +def test_resynthesis_rejects_invalid_parameters( |
| 583 | + parameter: str, value: int, message: str |
| 584 | +) -> None: |
| 585 | + hugr = from_coms(H(0), H(0)).to_python().modules[0] |
| 586 | + options: dict[str, Any] = {parameter: value} |
| 587 | + |
| 588 | + with pytest.raises(ValueError, match=f"^{message}$"): |
| 589 | + PauliGraphResynthesis(**options) |
| 590 | + |
| 591 | + optimisation = PauliGraphResynthesis() |
| 592 | + setattr(optimisation, parameter, value) |
| 593 | + |
| 594 | + with pytest.raises(ValueError, match=f"^{message}$"): |
| 595 | + optimisation.run(hugr, inplace=True) |
| 596 | + |
| 597 | + assert _count_hadamards(hugr) == 2 |
| 598 | + |
| 599 | + |
| 600 | +def test_wrapper_requires_enum() -> None: |
| 601 | + invalid_mode: Any = "on" |
| 602 | + with pytest.raises( |
| 603 | + TypeError, match="parallel_mode must be an instance of the ParallelMode enum" |
| 604 | + ): |
| 605 | + PauliGraphResynthesis(parallel_mode=invalid_mode) |
0 commit comments