microsoft/qdk

Public

mirrored from https://github.com/microsoft/qdkAvailable

CodeCommitsIssuesPull requestsActionsInsightsSecurity
copilot/add-link-to-qsharp-application

Branches

Tags

  • No tags available.
0Branches0Tags
Go to file
Add file
Code

Clone

HTTPS

Download ZIP

source/pip/tests/qre/test_enumeration.py

578lines · modecode

1# Copyright (c) Microsoft Corporation.
2# Licensed under the MIT License.
3
4from dataclasses import KW_ONLY, dataclass, field
5from enum import Enum
6from typing import cast
7
8import pytest
9
10from qsharp.qre import LOGICAL
11from qsharp.qre.models import SurfaceCode, GateBased, RoundBasedFactory
12from qsharp.qre.instruction_ids import LATTICE_SURGERY, T
13from qsharp.qre._isa_enumeration import (
14 ISARefNode,
15 _ComponentQuery,
16 _ProductNode,
17 _SumNode,
18)
19
20from .conftest import ExampleFactory, ExampleLogicalFactory
21
22
23def test_enumerate_instances():
24 """Test enumeration of SurfaceCode instances with default and custom domains."""
25 from qsharp.qre._enumeration import _enumerate_instances
26
27 instances = list(_enumerate_instances(SurfaceCode))
28
29 # There are 12 instances with distances from 3 to 25
30 assert len(instances) == 12
31 expected_distances = list(range(3, 26, 2))
32 for instance, expected_distance in zip(instances, expected_distances):
33 assert instance.distance == expected_distance
34
35 # Test with specific distances
36 instances = list(_enumerate_instances(SurfaceCode, distance=[3, 5, 7]))
37 assert len(instances) == 3
38 expected_distances = [3, 5, 7]
39 for instance, expected_distance in zip(instances, expected_distances):
40 assert instance.distance == expected_distance
41
42 # Test with fixed distance
43 instances = list(_enumerate_instances(SurfaceCode, distance=9))
44 assert len(instances) == 1
45 assert instances[0].distance == 9
46
47
48def test_enumerate_instances_bool():
49 """Test that boolean dataclass fields enumerate both True and False."""
50 from qsharp.qre._enumeration import _enumerate_instances
51
52 @dataclass
53 class BoolConfig:
54 _: KW_ONLY
55 flag: bool
56
57 instances = list(_enumerate_instances(BoolConfig))
58 assert len(instances) == 2
59 assert instances[0].flag is True
60 assert instances[1].flag is False
61
62
63def test_enumerate_instances_enum():
64 """Test that Enum dataclass fields enumerate all members."""
65 from qsharp.qre._enumeration import _enumerate_instances
66
67 class Color(Enum):
68 RED = 1
69 GREEN = 2
70 BLUE = 3
71
72 @dataclass
73 class EnumConfig:
74 _: KW_ONLY
75 color: Color
76
77 instances = list(_enumerate_instances(EnumConfig))
78 assert len(instances) == 3
79 assert instances[0].color == Color.RED
80 assert instances[1].color == Color.GREEN
81 assert instances[2].color == Color.BLUE
82
83
84def test_enumerate_instances_failure():
85 """Test that a field with no domain and no default raises ValueError."""
86 from qsharp.qre._enumeration import _enumerate_instances
87
88 @dataclass
89 class InvalidConfig:
90 _: KW_ONLY
91 # This field has no domain, is not bool/enum, and has no default
92 value: int
93
94 with pytest.raises(ValueError, match="Cannot enumerate field value"):
95 list(_enumerate_instances(InvalidConfig))
96
97
98def test_enumerate_instances_single():
99 """Test enumeration of a dataclass with a single non-kw-only field."""
100 from qsharp.qre._enumeration import _enumerate_instances
101
102 @dataclass
103 class SingleConfig:
104 value: int = 42
105
106 instances = list(_enumerate_instances(SingleConfig))
107 assert len(instances) == 1
108 assert instances[0].value == 42
109
110
111def test_enumerate_instances_literal():
112 """Test that Literal-typed fields enumerate their allowed values."""
113 from qsharp.qre._enumeration import _enumerate_instances
114
115 from typing import Literal
116
117 @dataclass
118 class LiteralConfig:
119 _: KW_ONLY
120 mode: Literal["fast", "slow"]
121
122 instances = list(_enumerate_instances(LiteralConfig))
123 assert len(instances) == 2
124 assert instances[0].mode == "fast"
125 assert instances[1].mode == "slow"
126
127
128def test_enumerate_instances_nested():
129 """Test enumeration of nested dataclass fields."""
130 from qsharp.qre._enumeration import _enumerate_instances
131
132 @dataclass
133 class InnerConfig:
134 _: KW_ONLY
135 option: bool
136
137 @dataclass
138 class OuterConfig:
139 _: KW_ONLY
140 inner: InnerConfig
141
142 instances = list(_enumerate_instances(OuterConfig))
143 assert len(instances) == 2
144 assert instances[0].inner.option is True
145 assert instances[1].inner.option is False
146
147
148def test_enumerate_instances_union():
149 """Test enumeration of union-typed dataclass fields."""
150 from qsharp.qre._enumeration import _enumerate_instances
151
152 @dataclass
153 class OptionA:
154 _: KW_ONLY
155 value: bool
156
157 @dataclass
158 class OptionB:
159 _: KW_ONLY
160 number: int = field(default=1, metadata={"domain": [1, 2, 3]})
161
162 @dataclass
163 class UnionConfig:
164 _: KW_ONLY
165 option: OptionA | OptionB
166
167 instances = list(_enumerate_instances(UnionConfig))
168 assert len(instances) == 5
169 assert isinstance(instances[0].option, OptionA)
170 assert instances[0].option.value is True
171 assert isinstance(instances[2].option, OptionB)
172 assert instances[2].option.number == 1
173
174
175def test_enumerate_instances_nested_with_constraints():
176 """Test constraining nested dataclass fields via a dict."""
177 from qsharp.qre._enumeration import _enumerate_instances
178
179 @dataclass
180 class InnerConfig:
181 _: KW_ONLY
182 option: bool
183
184 @dataclass
185 class OuterConfig:
186 _: KW_ONLY
187 inner: InnerConfig
188
189 # Constrain nested field via dict
190 instances = list(_enumerate_instances(OuterConfig, inner={"option": True}))
191 assert len(instances) == 1
192 assert instances[0].inner.option is True
193
194
195def test_enumerate_instances_union_single_type():
196 """Test restricting a union field to a single member type."""
197 from qsharp.qre._enumeration import _enumerate_instances
198
199 @dataclass
200 class OptionA:
201 _: KW_ONLY
202 value: bool
203
204 @dataclass
205 class OptionB:
206 _: KW_ONLY
207 number: int = field(default=1, metadata={"domain": [1, 2, 3]})
208
209 @dataclass
210 class UnionConfig:
211 _: KW_ONLY
212 option: OptionA | OptionB
213
214 # Restrict to OptionB only - uses its default domain
215 instances = list(_enumerate_instances(UnionConfig, option=OptionB))
216 assert len(instances) == 3
217 assert all(isinstance(i.option, OptionB) for i in instances)
218 assert [cast(OptionB, i.option).number for i in instances] == [1, 2, 3]
219
220 # Restrict to OptionA only
221 instances = list(_enumerate_instances(UnionConfig, option=OptionA))
222 assert len(instances) == 2
223 assert all(isinstance(i.option, OptionA) for i in instances)
224 assert cast(OptionA, instances[0].option).value is True
225 assert cast(OptionA, instances[1].option).value is False
226
227
228def test_enumerate_instances_union_list_of_types():
229 """Test restricting a union field to a subset of member types."""
230 from qsharp.qre._enumeration import _enumerate_instances
231
232 @dataclass
233 class OptionA:
234 _: KW_ONLY
235 value: bool
236
237 @dataclass
238 class OptionB:
239 _: KW_ONLY
240 number: int = field(default=1, metadata={"domain": [1, 2, 3]})
241
242 @dataclass
243 class OptionC:
244 _: KW_ONLY
245 flag: bool
246
247 @dataclass
248 class UnionConfig:
249 _: KW_ONLY
250 option: OptionA | OptionB | OptionC
251
252 # Select a subset: only OptionA and OptionB
253 instances = list(_enumerate_instances(UnionConfig, option=[OptionA, OptionB]))
254 assert len(instances) == 5 # 2 from OptionA + 3 from OptionB
255 assert all(isinstance(i.option, (OptionA, OptionB)) for i in instances)
256
257
258def test_enumerate_instances_union_constraint_dict():
259 """Test constraining union field members via a type-to-kwargs dict."""
260 from qsharp.qre._enumeration import _enumerate_instances
261
262 @dataclass
263 class OptionA:
264 _: KW_ONLY
265 value: bool
266
267 @dataclass
268 class OptionB:
269 _: KW_ONLY
270 number: int = field(default=1, metadata={"domain": [1, 2, 3]})
271
272 @dataclass
273 class UnionConfig:
274 _: KW_ONLY
275 option: OptionA | OptionB
276
277 # Constrain OptionA, enumerate only that member
278 instances = list(
279 _enumerate_instances(UnionConfig, option={OptionA: {"value": True}})
280 )
281 assert len(instances) == 1
282 assert isinstance(instances[0].option, OptionA)
283 assert instances[0].option.value is True
284
285 # Constrain OptionB with a domain, enumerate only that member
286 instances = list(
287 _enumerate_instances(UnionConfig, option={OptionB: {"number": [2, 3]}})
288 )
289 assert len(instances) == 2
290 assert all(isinstance(i.option, OptionB) for i in instances)
291 assert cast(OptionB, instances[0].option).number == 2
292 assert cast(OptionB, instances[1].option).number == 3
293
294 # Constrain one member and keep another with defaults
295 instances = list(
296 _enumerate_instances(
297 UnionConfig,
298 option={OptionA: {"value": True}, OptionB: {}},
299 )
300 )
301 assert len(instances) == 4 # 1 from OptionA + 3 from OptionB
302 assert isinstance(instances[0].option, OptionA)
303 assert instances[0].option.value is True
304 assert all(isinstance(i.option, OptionB) for i in instances[1:])
305 assert [cast(OptionB, i.option).number for i in instances[1:]] == [1, 2, 3]
306
307
308def test_enumerate_isas():
309 """Test ISA enumeration with products, sums, and hierarchical factories."""
310 ctx = GateBased(gate_time=50, measurement_time=100).context()
311
312 # This will enumerate the 4 ISAs for the error correction code
313 count = sum(1 for _ in SurfaceCode.q().enumerate(ctx))
314 assert count == 12
315
316 # This will enumerate the 2 ISAs for the error correction code when
317 # restricting the domain
318 count = sum(1 for _ in SurfaceCode.q(distance=[3, 4]).enumerate(ctx))
319 assert count == 2
320
321 # This will enumerate the 3 ISAs for the factory
322 count = sum(1 for _ in ExampleFactory.q().enumerate(ctx))
323 assert count == 3
324
325 # This will enumerate 36 ISAs for all products between the 12 error
326 # correction code ISAs and the 3 factory ISAs
327 count = sum(1 for _ in (SurfaceCode.q() * ExampleFactory.q()).enumerate(ctx))
328 assert count == 36
329
330 # When providing a list, components are chained (OR operation). This
331 # enumerates ISAs from first factory instance OR second factory instance
332 count = sum(
333 1
334 for _ in (
335 SurfaceCode.q() * (ExampleFactory.q() + ExampleFactory.q())
336 ).enumerate(ctx)
337 )
338 assert count == 72
339
340 # When providing separate arguments, components are combined via product
341 # (AND). This enumerates ISAs from first factory instance AND second
342 # factory instance
343 count = sum(
344 1
345 for _ in (SurfaceCode.q() * ExampleFactory.q() * ExampleFactory.q()).enumerate(
346 ctx
347 )
348 )
349 assert count == 108
350
351 # Hierarchical factory using from_components: the component receives ISAs
352 # from the product of other components as its source
353 count = sum(
354 1
355 for _ in (
356 SurfaceCode.q()
357 * ExampleLogicalFactory.q(source=(SurfaceCode.q() * ExampleFactory.q()))
358 ).enumerate(ctx)
359 )
360 assert count == 1296
361
362
363def test_binding_node():
364 """Test binding nodes with ISARefNode for component bindings"""
365 ctx = GateBased(gate_time=50, measurement_time=100).context()
366
367 # Test basic binding: same code used twice
368 # Without binding: 12 codes × 12 codes = 144 combinations
369 count_without = sum(1 for _ in (SurfaceCode.q() * SurfaceCode.q()).enumerate(ctx))
370 assert count_without == 144
371
372 # With binding: 12 codes (same instance used twice)
373 count_with = sum(
374 1
375 for _ in SurfaceCode.bind("c", ISARefNode("c") * ISARefNode("c")).enumerate(ctx)
376 )
377 assert count_with == 12
378
379 # Verify the binding works: with binding, both should use same params
380 for isa in SurfaceCode.bind("c", ISARefNode("c") * ISARefNode("c")).enumerate(ctx):
381 logical_gates = [g for g in isa if g.encoding == LOGICAL]
382 # Should have 1 logical gate (LATTICE_SURGERY)
383 assert len(logical_gates) == 1
384
385 # Test binding with factories (nested bindings)
386 count_without = sum(
387 1
388 for _ in (
389 SurfaceCode.q() * ExampleFactory.q() * SurfaceCode.q() * ExampleFactory.q()
390 ).enumerate(ctx)
391 )
392 assert count_without == 1296 # 12 * 3 * 12 * 3
393
394 count_with = sum(
395 1
396 for _ in SurfaceCode.bind(
397 "c",
398 ExampleFactory.bind(
399 "f",
400 ISARefNode("c") * ISARefNode("f") * ISARefNode("c") * ISARefNode("f"),
401 ),
402 ).enumerate(ctx)
403 )
404 assert count_with == 36 # 12 * 3
405
406 # Test binding with from_components equivalent (hierarchical)
407 # Without binding: 4 outer codes × (4 inner codes × 3 factories × 3 levels)
408 count_without = sum(
409 1
410 for _ in (
411 SurfaceCode.q()
412 * ExampleLogicalFactory.q(
413 source=(SurfaceCode.q() * ExampleFactory.q()),
414 )
415 ).enumerate(ctx)
416 )
417 assert count_without == 1296 # 12 * 12 * 3 * 3
418
419 # With binding: 4 codes (same used twice) × 3 factories × 3 levels
420 count_with = sum(
421 1
422 for _ in SurfaceCode.bind(
423 "c",
424 ISARefNode("c")
425 * ExampleLogicalFactory.q(
426 source=(ISARefNode("c") * ExampleFactory.q()),
427 ),
428 ).enumerate(ctx)
429 )
430 assert count_with == 108 # 12 * 3 * 3
431
432 # Test binding with kwargs
433 count_with_kwargs = sum(
434 1
435 for _ in SurfaceCode.q(distance=5)
436 .bind("c", ISARefNode("c") * ISARefNode("c"))
437 .enumerate(ctx)
438 )
439 assert count_with_kwargs == 1 # Only distance=5
440
441 # Verify kwargs are applied
442 for isa in (
443 SurfaceCode.q(distance=5)
444 .bind("c", ISARefNode("c") * ISARefNode("c"))
445 .enumerate(ctx)
446 ):
447 logical_gates = [g for g in isa if g.encoding == LOGICAL]
448 assert all(g.space(1) == 49 for g in logical_gates)
449
450 # Test multiple independent bindings (nested)
451 count = sum(
452 1
453 for _ in SurfaceCode.bind(
454 "c1",
455 ExampleFactory.bind(
456 "c2",
457 ISARefNode("c1")
458 * ISARefNode("c1")
459 * ISARefNode("c2")
460 * ISARefNode("c2"),
461 ),
462 ).enumerate(ctx)
463 )
464 # 12 codes for c1 × 3 factories for c2
465 assert count == 36
466
467
468def test_binding_node_errors():
469 """Test error handling for binding nodes"""
470 ctx = GateBased(gate_time=50, measurement_time=100).context()
471
472 # Test ISARefNode enumerate with undefined binding raises ValueError
473 try:
474 list(ISARefNode("test").enumerate(ctx))
475 assert False, "Should have raised ValueError"
476 except ValueError as e:
477 assert "Undefined component reference: 'test'" in str(e)
478
479
480def test_product_isa_enumeration_nodes():
481 """Test that multiplying ISAQuery nodes produces flattened ProductNodes."""
482 terminal = SurfaceCode.q()
483 query = terminal * terminal
484
485 # Multiplication should create ProductNode
486 assert isinstance(query, _ProductNode)
487 assert len(query.sources) == 2
488 for source in query.sources:
489 assert isinstance(source, _ComponentQuery)
490
491 # Multiplying again should extend the sources
492 query = query * terminal
493 assert isinstance(query, _ProductNode)
494 assert len(query.sources) == 3
495 for source in query.sources:
496 assert isinstance(source, _ComponentQuery)
497
498 # Also from the other side
499 query = terminal * query
500 assert isinstance(query, _ProductNode)
501 assert len(query.sources) == 4
502 for source in query.sources:
503 assert isinstance(source, _ComponentQuery)
504
505 # Also for two ProductNodes
506 query = query * query
507 assert isinstance(query, _ProductNode)
508 assert len(query.sources) == 8
509 for source in query.sources:
510 assert isinstance(source, _ComponentQuery)
511
512
513def test_sum_isa_enumeration_nodes():
514 """Test that adding ISAQuery nodes produces flattened SumNodes."""
515 terminal = SurfaceCode.q()
516 query = terminal + terminal
517
518 # Multiplication should create SumNode
519 assert isinstance(query, _SumNode)
520 assert len(query.sources) == 2
521 for source in query.sources:
522 assert isinstance(source, _ComponentQuery)
523
524 # Multiplying again should extend the sources
525 query = query + terminal
526 assert isinstance(query, _SumNode)
527 assert len(query.sources) == 3
528 for source in query.sources:
529 assert isinstance(source, _ComponentQuery)
530
531 # Also from the other side
532 query = terminal + query
533 assert isinstance(query, _SumNode)
534 assert len(query.sources) == 4
535 for source in query.sources:
536 assert isinstance(source, _ComponentQuery)
537
538 # Also for two SumNodes
539 query = query + query
540 assert isinstance(query, _SumNode)
541 assert len(query.sources) == 8
542 for source in query.sources:
543 assert isinstance(source, _ComponentQuery)
544
545
546def test_round_based_model_does_not_expose_code_query_instructions():
547 """
548 Tests that the round-based model does not expose instructions from the inner
549 code query.
550 """
551
552 arch = GateBased(gate_time=100, measurement_time=500)
553 ctx = arch.context()
554 RoundBasedFactory.q(use_cache=False).populate(ctx)
555
556 graph = ctx._provenance
557
558 # Accumulate total space and time for all generated factories to make sure
559 # they are unchanged
560 total_space, total_time = 0, 0
561
562 for node in range(1, graph.num_nodes() + 1):
563 instruction = graph.instruction(node)
564 if instruction.id == LATTICE_SURGERY:
565 assert (
566 False
567 ), "Lattice surgery instruction found in round-based model provenance graph: {instruction}"
568
569 if instruction.id == T:
570 total_space += instruction.expect_space()
571 total_time += instruction.expect_time()
572
573 assert total_time == 34_375_000
574 assert total_space == 12_946_489
575
576 assert (
577 True
578 ), "Lattice surgery instruction not found in round-based model provenance graph"
579