diff --git a/src/pyqasm/pulse/validator.py b/src/pyqasm/pulse/validator.py index 888f818..2c95764 100644 --- a/src/pyqasm/pulse/validator.py +++ b/src/pyqasm/pulse/validator.py @@ -114,11 +114,10 @@ def validate_duration_or_stretch_statements( Generic validation function for DurationType and StretchType declarations or assignments. Args: - statement: The AST statement node - statement_type: The expected AST node type - base_type: The declared type (DurationType or StretchType) - rvalue: The initializer or assigned value - global_scope: Global symbol table. + statement (Statement): The AST statement node. + base_type (Any): The declared type, function does nothing if not DurationType or StretchType. + rvalue (Any): The initializer or assigned value. + global_scope (dict): Global symbol table. Raises: ValidationError: If the assigned value is not a DurationLiteral, diff --git a/src/pyqasm/visitor.py b/src/pyqasm/visitor.py index 62f52fa..7b2b90a 100644 --- a/src/pyqasm/visitor.py +++ b/src/pyqasm/visitor.py @@ -20,13 +20,14 @@ """ import copy +import functools import logging import re import sys from collections import OrderedDict, deque from functools import partial from io import StringIO -from typing import Any, Callable, Optional, Sequence, cast +from typing import Any, Callable, Optional, Sequence, TypeVar, cast import numpy as np import openqasm3.ast as qasm3_ast @@ -83,6 +84,22 @@ logger = logging.getLogger(__name__) logger.propagate = False +F = TypeVar("F", bound=Callable[..., Any]) + + +def check_only_return_empty(func: F) -> F: + """Decorator for functions which use check_only to return an empty list.""" + + @functools.wraps(func) + def wrapper(self, *args, **kwargs): + """Wrapper that intercepts the return value and replaces it with an empty list.""" + result = func(self, *args, **kwargs) + if self._check_only: + return [] + return result + + return wrapper # type: ignore + # pylint: disable-next=too-many-instance-attributes class QasmVisitor: @@ -199,6 +216,7 @@ def _construct_visit_map(self): qasm3_ast.CalibrationGrammarDeclaration: self._visit_calibration_grammar_declaration, } + @check_only_return_empty def _visit_quantum_register( self, register: qasm3_ast.QubitDeclaration ) -> list[qasm3_ast.QubitDeclaration]: @@ -280,8 +298,6 @@ def _visit_quantum_register( logger.debug("Added labels for register '%s'", str(register)) - if self._check_only: - return [] return [register] # pylint: disable-next=too-many-locals,too-many-branches,too-many-statements @@ -530,26 +546,6 @@ def _qubit_register_consolidation( return _valid_statements - def _handle_function_init_expression( - self, expression: qasm3_ast.FunctionCall, init_value: Any - ) -> None | qasm3_ast.Expression: - """Handle function initialization expression. - - Args: - expression (FunctionCall): The statement to handle function initialization expression. - init_value (Any): The value to handle function initialization expression. - - Returns: - None | Expression: The resultant expression if - the expression is applied, otherwise None. - """ - if isinstance(expression, qasm3_ast.FunctionCall): - func_name = expression.name.name - if func_name in FUNCTION_MAP: - if isinstance(init_value, (float, int)): - return qasm3_ast.FloatLiteral(init_value) - return None - def _handle_extern_function_cleanup( self, statements: list, statement: qasm3_ast.Statement ) -> None: @@ -718,7 +714,6 @@ def _visit_measurement( # pylint: disable=too-many-locals,too-many-branches,too if self._check_only: return [] - return unrolled_measurements def _resolve_unindexed_reset_qubit(self, statement: qasm3_ast.QuantumReset) -> bool: @@ -762,12 +757,13 @@ def _visit_reset(self, statement: qasm3_ast.QuantumReset) -> list[qasm3_ast.Quan or an empty list if self._check_only is True. """ logger.debug("Visiting reset statement '%s'", str(statement)) + if self._resolve_unindexed_reset_qubit(statement): return [statement] - if len(self._function_qreg_size_map) > 0: # atleast in SOME function scope + if len(self._function_qreg_size_map) > 0: # at least in SOME function scope # since we may have multiple function scopes, we need to transform the qubits - # to use the global qreg identifiers + # to use the global qreg identifiers. for transform_map, size_map in zip( reversed(self._function_qreg_transform_map), reversed(self._function_qreg_size_map) ): @@ -807,7 +803,6 @@ def _visit_reset(self, statement: qasm3_ast.QuantumReset) -> list[qasm3_ast.Quan if self._check_only: return [] - return unrolled_resets def _expand_barrier_ranges( @@ -1137,6 +1132,7 @@ def _update_qubit_depth_for_gate( qubit_node.depth = max_involved_depth # pylint: disable=too-many-branches, too-many-locals + @check_only_return_empty def _visit_basic_gate_operation( self, operation: qasm3_ast.QuantumGate, @@ -1235,9 +1231,6 @@ def _visit_basic_gate_operation( for final_gate in result: Qasm3Analyzer.verify_gate_qubits(final_gate, operation.span) - if self._check_only: - return [] - return result def _visit_break(self, statement: qasm3_ast.BreakStatement) -> None: @@ -1252,6 +1245,7 @@ def _visit_continue(self, statement: qasm3_ast.ContinueStatement) -> None: error_node=statement, ) + @check_only_return_empty def _visit_custom_gate_operation( self, operation: qasm3_ast.QuantumGate, @@ -1352,11 +1346,9 @@ def _visit_custom_gate_operation( self._scope_manager.pop_scope() self._scope_manager.restore_context() - if self._check_only: - return [] - return result + @check_only_return_empty def _visit_external_gate_operation( self, operation: qasm3_ast.QuantumGate, @@ -1428,11 +1420,10 @@ def gate_function(*qubits): Qasm3Analyzer.verify_gate_qubits(final_gate, operation.span) self._scope_manager.restore_context() - if self._check_only: - return [] return result + @check_only_return_empty def _visit_phase_operation( self, operation: qasm3_ast.QuantumPhase, @@ -1492,9 +1483,6 @@ def _visit_phase_operation( # if it were in function scope, then the args would have been evaluated and added to the # qubit list - if self._check_only: - return [] - return [operation] def _visit_generic_gate_operation( # pylint: disable=too-many-branches, too-many-statements @@ -1659,9 +1647,9 @@ def _visit_generic_gate_operation( # pylint: disable=too-many-branches, too-man if self._check_only: return [] - return result + @check_only_return_empty def _visit_constant_declaration( self, statement: qasm3_ast.ConstantDeclaration ) -> list[qasm3_ast.Statement]: @@ -1772,18 +1760,16 @@ def _visit_constant_declaration( statement.init_expression = PulseValidator.make_complex_binary_expression(init_value) if isinstance(statement.init_expression, qasm3_ast.FunctionCall): - statement.init_expression = ( - self._handle_function_init_expression(statement.init_expression, init_value) - or statement.init_expression - ) - self._handle_extern_function_cleanup(statements, statement) + function_name = statement.init_expression.name.name + if function_name in FUNCTION_MAP and isinstance(init_value, (float, int)): + statement.init_expression = qasm3_ast.FloatLiteral(init_value) - if self._check_only: - return [] + self._handle_extern_function_cleanup(statements, statement) return statements # pylint: disable=too-many-branches, too-many-statements, too-many-locals + @check_only_return_empty def _visit_classical_declaration( self, statement: qasm3_ast.ClassicalDeclaration ) -> list[qasm3_ast.Statement]: @@ -1883,7 +1869,6 @@ def _visit_classical_declaration( # populate the variable if statement.init_expression: - global_scope = self._scope_manager.get_global_scope() PulseValidator.validate_duration_or_stretch_statements( statement=statement, @@ -2013,16 +1998,13 @@ def _visit_classical_declaration( statement.init_expression = PulseValidator.make_complex_binary_expression(init_value) if isinstance(statement.init_expression, qasm3_ast.FunctionCall): - statement.init_expression = ( - self._handle_function_init_expression(statement.init_expression, init_value) - or statement.init_expression - ) - - if self._check_only: - return [] + function_name = statement.init_expression.name.name + if function_name in FUNCTION_MAP and isinstance(init_value, (float, int)): + statement.init_expression = qasm3_ast.FloatLiteral(init_value) return statements + @check_only_return_empty def _visit_classical_assignment( self, statement: qasm3_ast.ClassicalAssignment ) -> list[qasm3_ast.Statement]: @@ -2187,16 +2169,12 @@ def _visit_classical_assignment( ) if isinstance(statement.rvalue, qasm3_ast.FunctionCall): - statement.rvalue = ( - self._handle_function_init_expression(statement.rvalue, rvalue_eval) - or statement.rvalue - ) + function_name = statement.rvalue.name.name + if function_name in FUNCTION_MAP and isinstance(rvalue_eval, (float, int)): + statement.rvalue = qasm3_ast.FloatLiteral(rvalue_eval) self._handle_extern_function_cleanup(statements, statement) - if self._check_only: - return [] - return statements def _evaluate_array_initialization( @@ -2241,6 +2219,7 @@ def _update_branching_gate_depths(self) -> None: self._is_branch_clbits.clear() self._is_branch_qubits.clear() + @check_only_return_empty def _visit_branching_statement( self, statement: qasm3_ast.BranchingStatement ) -> list[qasm3_ast.Statement]: @@ -2370,9 +2349,6 @@ def ravel(bit_ind): if not self._in_branching_statement: self._update_branching_gate_depths() - if self._check_only: - return [] - return result # type: ignore[return-value] def _visit_forin_loop(self, statement: qasm3_ast.ForInLoop) -> list[qasm3_ast.Statement]: @@ -2459,6 +2435,7 @@ def _visit_forin_loop(self, statement: qasm3_ast.ForInLoop) -> list[qasm3_ast.St return [] return result + @check_only_return_empty def _visit_subroutine_definition( self, statement: qasm3_ast.SubroutineDefinition | qasm3_ast.ExternDeclaration ) -> Sequence[None | qasm3_ast.ExternDeclaration]: @@ -2508,8 +2485,6 @@ def _visit_subroutine_definition( statements.append(statement) self._subroutine_defns[fn_name] = statement - if self._check_only: - return [] return statements @@ -2724,7 +2699,7 @@ def _visit_alias_statement(self, statement: qasm3_ast.AliasStatement) -> list[No # this will only build a global alias map - # whenever we are referring to qubits , we will first check in the global map of registers + # whenever we are referring to qubits, we will first check in the global map of registers # if the register is present, we will use the global map to get the qubit labels # if not, we will check the alias map for the labels @@ -2873,8 +2848,6 @@ def _evaluate_case(statements): self._scope_manager.pop_scope() self._scope_manager.restore_context() - if self._check_only: - return [] return result case_fulfilled = False @@ -3384,6 +3357,7 @@ def _visit_calibration_grammar_declaration( return [statement] + @check_only_return_empty def _visit_include(self, include: qasm3_ast.Include) -> list[qasm3_ast.Statement]: """Visit an include statement element. @@ -3400,8 +3374,6 @@ def _visit_include(self, include: qasm3_ast.Include) -> list[qasm3_ast.Statement f"File '{filename}' already included", error_node=include, span=include.span ) self._included_files.add(filename) - if self._check_only: - return [] return [include]