Download python/qlambda/compiler.py from Snapkitty/topological-quantum-computer: direct link, hf CLI and curl.
- Browser
- Download file 19.6 kB
-
https://huggingface.co/Snapkitty/topological-quantum-computer/resolve/main/python/qlambda/compiler.py
- Command line
-
hf download hf://Snapkitty/topological-quantum-computer/python/qlambda/compiler.py
-
curl -L -o compiler.py https://huggingface.co/Snapkitty/topological-quantum-computer/resolve/main/python/qlambda/compiler.py
19.6 kB
| """Q-Lambda lexer, parser, reversible QIR synthesizer, and uncompute pass.""" | |
| from __future__ import annotations | |
| from dataclasses import dataclass | |
| from enum import Enum | |
| import re | |
| from typing import Dict, Iterable, List, Optional, Tuple, Union | |
| class TokenType(Enum): | |
| KW_QBIT = "qbit" | |
| KW_LET = "let" | |
| KW_REVERSED = "reversed" | |
| KW_WITH = "with" | |
| KW_DO = "do" | |
| KW_ORACLE = "oracle" | |
| OP_ROTR = ">>>" | |
| OP_SHR = ">>" | |
| OP_XOR = "^" | |
| OP_AND = "&" | |
| OP_OR = "|" | |
| OP_NOT = "~" | |
| OP_ASSIGN = "=" | |
| OP_ADD = "+" | |
| LPAREN = "(" | |
| RPAREN = ")" | |
| LBRACE = "{" | |
| RBRACE = "}" | |
| LBRACK = "[" | |
| RBRACK = "]" | |
| COMMA = "," | |
| COLON = ":" | |
| SEMI = ";" | |
| IDENT = "IDENT" | |
| INTEGER = "INTEGER" | |
| EOF = "EOF" | |
| class Token: | |
| type: TokenType | |
| value: str | |
| line: int | |
| col: int | |
| class Lexer: | |
| TOKEN_REGEX: Tuple[Tuple[TokenType, str], ...] = ( | |
| (TokenType.KW_QBIT, r"\bqbit\b"), | |
| (TokenType.KW_LET, r"\blet\b"), | |
| (TokenType.KW_REVERSED, r"\breversed\b"), | |
| (TokenType.KW_WITH, r"\bwith\b"), | |
| (TokenType.KW_DO, r"\bdo\b"), | |
| (TokenType.KW_ORACLE, r"\boracle\b"), | |
| (TokenType.OP_ROTR, r">>>"), | |
| (TokenType.OP_SHR, r">>"), | |
| (TokenType.OP_XOR, r"\^"), | |
| (TokenType.OP_AND, r"&"), | |
| (TokenType.OP_OR, r"\|"), | |
| (TokenType.OP_NOT, r"~"), | |
| (TokenType.OP_ASSIGN, r"="), | |
| (TokenType.OP_ADD, r"\+"), | |
| (TokenType.LPAREN, r"\("), | |
| (TokenType.RPAREN, r"\)"), | |
| (TokenType.LBRACE, r"\{"), | |
| (TokenType.RBRACE, r"\}"), | |
| (TokenType.LBRACK, r"\["), | |
| (TokenType.RBRACK, r"\]"), | |
| (TokenType.COMMA, r","), | |
| (TokenType.COLON, r":"), | |
| (TokenType.SEMI, r";"), | |
| (TokenType.INTEGER, r"\b\d+\b"), | |
| (TokenType.IDENT, r"[a-zA-Z_][a-zA-Z0-9_]*"), | |
| ) | |
| def __init__(self, source: str): | |
| self.source = re.sub(r"//.*", "", expand_unroll_directives(source)) | |
| self.pos = 0 | |
| self.line = 1 | |
| self.col = 1 | |
| def tokenize(self) -> List[Token]: | |
| tokens: List[Token] = [] | |
| compiled = [(kind, re.compile(pattern)) for kind, pattern in self.TOKEN_REGEX] | |
| while self.pos < len(self.source): | |
| char = self.source[self.pos] | |
| if char == "\n": | |
| self.line += 1 | |
| self.col = 1 | |
| self.pos += 1 | |
| continue | |
| if char.isspace(): | |
| self.col += 1 | |
| self.pos += 1 | |
| continue | |
| for tok_type, regex in compiled: | |
| match = regex.match(self.source, self.pos) | |
| if match: | |
| value = match.group(0) | |
| tokens.append(Token(tok_type, value, self.line, self.col)) | |
| self.pos += len(value) | |
| self.col += len(value) | |
| break | |
| else: | |
| raise SyntaxError(f"Unexpected character {char!r} at line {self.line}, col {self.col}") | |
| tokens.append(Token(TokenType.EOF, "", self.line, self.col)) | |
| return tokens | |
| class ASTNode: | |
| """Base class for Q-Lambda AST nodes.""" | |
| class TypeNode(ASTNode): | |
| name: str | |
| size: int | |
| class ExprNode(ASTNode): | |
| """Base class for Q-Lambda expressions.""" | |
| class VarExpr(ExprNode): | |
| name: str | |
| class IntLiteralExpr(ExprNode): | |
| value: int | |
| class BinOpExpr(ExprNode): | |
| op: TokenType | |
| left: ExprNode | |
| right: ExprNode | |
| class UnOpExpr(ExprNode): | |
| op: TokenType | |
| operand: ExprNode | |
| class StmtNode(ASTNode): | |
| """Base class for Q-Lambda statements.""" | |
| class LetStmt(StmtNode): | |
| var_name: str | |
| var_type: TypeNode | |
| value: ExprNode | |
| class WithDoStmt(StmtNode): | |
| bindings: List[LetStmt] | |
| body: List[StmtNode] | |
| class ReversedBlockStmt(StmtNode): | |
| body: List[StmtNode] | |
| class OracleDeclStmt(StmtNode): | |
| name: str | |
| params: List[Tuple[str, TypeNode]] | |
| returns: TypeNode | |
| body: List[StmtNode] | |
| class Parser: | |
| PRECEDENCE = { | |
| TokenType.OP_XOR: 1, | |
| TokenType.OP_OR: 1, | |
| TokenType.OP_AND: 2, | |
| TokenType.OP_ADD: 2, | |
| TokenType.OP_ROTR: 3, | |
| TokenType.OP_SHR: 3, | |
| } | |
| def __init__(self, tokens: List[Token]): | |
| self.tokens = tokens | |
| self.pos = 0 | |
| def peek(self) -> Token: | |
| return self.tokens[self.pos] | |
| def consume(self, expected_type: TokenType) -> Token: | |
| current = self.peek() | |
| if current.type != expected_type: | |
| raise SyntaxError( | |
| f"Expected {expected_type.value}, got {current.type.value} " | |
| f"({current.value!r}) at line {current.line}" | |
| ) | |
| self.pos += 1 | |
| return current | |
| def parse(self) -> List[StmtNode]: | |
| nodes: List[StmtNode] = [] | |
| while self.peek().type != TokenType.EOF: | |
| nodes.append(self.parse_statement()) | |
| return nodes | |
| def parse_statement(self) -> StmtNode: | |
| token = self.peek() | |
| if token.type == TokenType.KW_LET: | |
| return self.parse_let() | |
| if token.type == TokenType.KW_WITH: | |
| return self.parse_with_do() | |
| if token.type == TokenType.KW_REVERSED: | |
| return self.parse_reversed() | |
| if token.type == TokenType.KW_ORACLE: | |
| return self.parse_oracle() | |
| raise SyntaxError(f"Unexpected statement starting with {token.value!r}") | |
| def parse_let(self) -> LetStmt: | |
| self.consume(TokenType.KW_LET) | |
| var_name = self.consume(TokenType.IDENT).value | |
| self.consume(TokenType.COLON) | |
| var_type = self.parse_type() | |
| self.consume(TokenType.OP_ASSIGN) | |
| value = self.parse_expr() | |
| self.consume(TokenType.SEMI) | |
| return LetStmt(var_name, var_type, value) | |
| def parse_type(self) -> TypeNode: | |
| type_name = self.consume(TokenType.KW_QBIT).value | |
| self.consume(TokenType.LBRACK) | |
| size = int(self.consume(TokenType.INTEGER).value) | |
| self.consume(TokenType.RBRACK) | |
| return TypeNode(type_name, size) | |
| def parse_with_do(self) -> WithDoStmt: | |
| self.consume(TokenType.KW_WITH) | |
| self.consume(TokenType.LPAREN) | |
| bindings: List[LetStmt] = [] | |
| while self.peek().type != TokenType.RPAREN: | |
| bindings.append(self.parse_let()) | |
| self.consume(TokenType.RPAREN) | |
| self.consume(TokenType.KW_DO) | |
| body = self.parse_block() | |
| return WithDoStmt(bindings, body) | |
| def parse_reversed(self) -> ReversedBlockStmt: | |
| self.consume(TokenType.KW_REVERSED) | |
| return ReversedBlockStmt(self.parse_block()) | |
| def parse_oracle(self) -> OracleDeclStmt: | |
| self.consume(TokenType.KW_ORACLE) | |
| name = self.consume(TokenType.IDENT).value | |
| self.consume(TokenType.LPAREN) | |
| params: List[Tuple[str, TypeNode]] = [] | |
| while self.peek().type != TokenType.RPAREN: | |
| param_name = self.consume(TokenType.IDENT).value | |
| self.consume(TokenType.COLON) | |
| params.append((param_name, self.parse_type())) | |
| if self.peek().type == TokenType.COMMA: | |
| self.consume(TokenType.COMMA) | |
| self.consume(TokenType.RPAREN) | |
| self.consume(TokenType.COLON) | |
| returns = self.parse_type() | |
| body = self.parse_block() | |
| return OracleDeclStmt(name, params, returns, body) | |
| def parse_block(self) -> List[StmtNode]: | |
| self.consume(TokenType.LBRACE) | |
| body: List[StmtNode] = [] | |
| while self.peek().type != TokenType.RBRACE: | |
| body.append(self.parse_statement()) | |
| self.consume(TokenType.RBRACE) | |
| return body | |
| def parse_expr(self) -> ExprNode: | |
| return self.parse_binary_expr(0) | |
| def parse_binary_expr(self, precedence: int) -> ExprNode: | |
| left = self.parse_primary() | |
| while True: | |
| op = self.peek().type | |
| if op not in self.PRECEDENCE or self.PRECEDENCE[op] < precedence: | |
| break | |
| self.consume(op) | |
| right = self.parse_binary_expr(self.PRECEDENCE[op] + 1) | |
| left = BinOpExpr(op, left, right) | |
| return left | |
| def parse_primary(self) -> ExprNode: | |
| token = self.peek() | |
| if token.type == TokenType.OP_NOT: | |
| self.consume(TokenType.OP_NOT) | |
| return UnOpExpr(TokenType.OP_NOT, self.parse_primary()) | |
| if token.type == TokenType.IDENT: | |
| return VarExpr(self.consume(TokenType.IDENT).value) | |
| if token.type == TokenType.INTEGER: | |
| return IntLiteralExpr(int(self.consume(TokenType.INTEGER).value)) | |
| if token.type == TokenType.LPAREN: | |
| self.consume(TokenType.LPAREN) | |
| expr = self.parse_expr() | |
| self.consume(TokenType.RPAREN) | |
| return expr | |
| raise SyntaxError(f"Unexpected expression token: {token.value!r}") | |
| class QIRInstruction: | |
| gate: str | |
| controls: Tuple[int, ...] | |
| targets: Tuple[int, ...] | |
| params: Tuple[Union[int, float], ...] = () | |
| class QIREngine: | |
| def __init__(self): | |
| self.qubit_counter = 0 | |
| self.instructions: List[QIRInstruction] = [] | |
| self.scopes: List[Dict[str, List[int]]] = [{}] | |
| def allocate(self, name: str, size: int) -> List[int]: | |
| if size <= 0: | |
| raise ValueError("qbit register size must be positive") | |
| qids = list(range(self.qubit_counter, self.qubit_counter + size)) | |
| self.qubit_counter += size | |
| self.scopes[-1][name] = qids | |
| return qids | |
| def resolve(self, name: str) -> List[int]: | |
| for scope in reversed(self.scopes): | |
| if name in scope: | |
| return scope[name] | |
| raise NameError(f"Quantum register {name!r} not found") | |
| def emit( | |
| self, | |
| gate: str, | |
| controls: Iterable[int] = (), | |
| targets: Iterable[int] = (), | |
| params: Iterable[Union[int, float]] = (), | |
| ) -> None: | |
| self.instructions.append( | |
| QIRInstruction(gate, tuple(controls), tuple(targets), tuple(params)) | |
| ) | |
| def push_scope(self) -> None: | |
| self.scopes.append({}) | |
| def pop_scope(self) -> Dict[str, List[int]]: | |
| if len(self.scopes) == 1: | |
| raise RuntimeError("cannot pop root scope") | |
| return self.scopes.pop() | |
| class QLambdaCompiler: | |
| def __init__(self, ast: List[StmtNode]): | |
| self.ast = ast | |
| self.qir = QIREngine() | |
| def compile(self) -> List[QIRInstruction]: | |
| for node in self.ast: | |
| self.visit(node) | |
| return self.qir.instructions | |
| def visit(self, node: StmtNode) -> None: | |
| if isinstance(node, OracleDeclStmt): | |
| self.visit_oracle(node) | |
| elif isinstance(node, LetStmt): | |
| self.visit_let(node) | |
| elif isinstance(node, WithDoStmt): | |
| self.visit_with_do(node) | |
| elif isinstance(node, ReversedBlockStmt): | |
| self.visit_reversed(node) | |
| else: | |
| raise NotImplementedError(type(node).__name__) | |
| def visit_oracle(self, node: OracleDeclStmt) -> None: | |
| self.qir.push_scope() | |
| for param_name, param_type in node.params: | |
| self.qir.allocate(param_name, param_type.size) | |
| self.qir.allocate(f"{node.name}_out", node.returns.size) | |
| for stmt in node.body: | |
| self.visit(stmt) | |
| self.qir.pop_scope() | |
| def visit_let(self, node: LetStmt) -> None: | |
| target = self.qir.allocate(node.var_name, node.var_type.size) | |
| self.synthesize_expr(node.value, target) | |
| def visit_with_do(self, node: WithDoStmt) -> None: | |
| self.qir.push_scope() | |
| binding_start = len(self.qir.instructions) | |
| for binding in node.bindings: | |
| self.visit_let(binding) | |
| binding_end = len(self.qir.instructions) | |
| for stmt in node.body: | |
| self.visit(stmt) | |
| self.qir.instructions.extend(self.invert_circuit(self.qir.instructions[binding_start:binding_end])) | |
| self.qir.pop_scope() | |
| def visit_reversed(self, node: ReversedBlockStmt) -> None: | |
| marker = len(self.qir.instructions) | |
| for stmt in node.body: | |
| self.visit(stmt) | |
| self.qir.instructions = self.qir.instructions[:marker] + self.invert_circuit( | |
| self.qir.instructions[marker:] | |
| ) | |
| def synthesize_expr(self, expr: ExprNode, target: List[int]) -> None: | |
| if isinstance(expr, VarExpr): | |
| self.copy_register(self.qir.resolve(expr.name), target) | |
| elif isinstance(expr, IntLiteralExpr): | |
| for index, qid in enumerate(target): | |
| if (expr.value >> index) & 1: | |
| self.qir.emit("X", targets=[qid]) | |
| elif isinstance(expr, UnOpExpr) and expr.op == TokenType.OP_NOT: | |
| self.synthesize_expr(expr.operand, target) | |
| for qid in target: | |
| self.qir.emit("X", targets=[qid]) | |
| elif isinstance(expr, BinOpExpr): | |
| self.synthesize_binop(expr, target) | |
| else: | |
| raise NotImplementedError(f"Cannot synthesize {expr!r}") | |
| def synthesize_binop(self, expr: BinOpExpr, target: List[int]) -> None: | |
| if expr.op == TokenType.OP_XOR: | |
| self.synthesize_expr(expr.left, target) | |
| self.synthesize_expr(expr.right, target) | |
| elif expr.op == TokenType.OP_AND: | |
| marker = len(self.qir.instructions) | |
| left = self.temp("_and_l", len(target)) | |
| right = self.temp("_and_r", len(target)) | |
| self.synthesize_expr(expr.left, left) | |
| self.synthesize_expr(expr.right, right) | |
| temp_program = self.qir.instructions[marker:] | |
| for lq, rq, tq in zip(left, right, target): | |
| self.qir.emit("CCX", controls=[lq, rq], targets=[tq]) | |
| self.qir.instructions.extend(self.invert_circuit(temp_program)) | |
| elif expr.op == TokenType.OP_ROTR: | |
| shift = self.literal_shift(expr.right) | |
| src = self.temp("_rotr", len(target)) | |
| self.synthesize_expr(expr.left, src) | |
| width = len(target) | |
| for index, tq in enumerate(target): | |
| self.qir.emit("CX", controls=[src[(index + shift) % width]], targets=[tq]) | |
| elif expr.op == TokenType.OP_SHR: | |
| shift = self.literal_shift(expr.right) | |
| src = self.temp("_shr", len(target)) | |
| self.synthesize_expr(expr.left, src) | |
| for index, tq in enumerate(target): | |
| src_index = index + shift | |
| if src_index < len(src): | |
| self.qir.emit("CX", controls=[src[src_index]], targets=[tq]) | |
| elif expr.op == TokenType.OP_ADD: | |
| self.synthesize_modular_add(expr.left, expr.right, target) | |
| else: | |
| raise NotImplementedError(f"Unsupported operator {expr.op.value}") | |
| def synthesize_modular_add(self, left_expr: ExprNode, right_expr: ExprNode, target: List[int]) -> None: | |
| width = len(target) | |
| left = self.temp("_add_l", width) | |
| right = self.temp("_add_r", width) | |
| carry = self.temp("_carry", width + 1) | |
| self.synthesize_expr(left_expr, left) | |
| self.synthesize_expr(right_expr, right) | |
| self.copy_register(left, target) | |
| for index in range(width): | |
| self.qir.emit("CCX", controls=[target[index], right[index]], targets=[carry[index + 1]]) | |
| self.qir.emit("CX", controls=[target[index]], targets=[right[index]]) | |
| self.qir.emit("CCX", controls=[right[index], carry[index]], targets=[carry[index + 1]]) | |
| self.qir.emit("CX", controls=[right[index]], targets=[target[index]]) | |
| for index in range(width - 1, -1, -1): | |
| self.qir.emit("CX", controls=[right[index]], targets=[target[index]]) | |
| self.qir.emit("CCX", controls=[right[index], carry[index]], targets=[carry[index + 1]]) | |
| self.qir.emit("CX", controls=[target[index]], targets=[right[index]]) | |
| self.qir.emit("CCX", controls=[target[index], right[index]], targets=[carry[index + 1]]) | |
| def copy_register(self, source: List[int], target: List[int]) -> None: | |
| if len(source) < len(target): | |
| raise ValueError("source register is narrower than target register") | |
| for src, dst in zip(source, target): | |
| self.qir.emit("CX", controls=[src], targets=[dst]) | |
| def temp(self, prefix: str, size: int) -> List[int]: | |
| return self.qir.allocate(f"{prefix}_{self.qir.qubit_counter}", size) | |
| def literal_shift(expr: ExprNode) -> int: | |
| if not isinstance(expr, IntLiteralExpr): | |
| raise TypeError("shift/rotate amount must be an integer literal") | |
| return expr.value | |
| def invert_circuit(instructions: List[QIRInstruction]) -> List[QIRInstruction]: | |
| inverted: List[QIRInstruction] = [] | |
| for inst in reversed(instructions): | |
| if inst.gate in {"X", "CX", "CCX", "H"}: | |
| inverted.append(inst) | |
| elif inst.gate == "T": | |
| inverted.append(QIRInstruction("TDG", inst.controls, inst.targets, inst.params)) | |
| elif inst.gate == "TDG": | |
| inverted.append(QIRInstruction("T", inst.controls, inst.targets, inst.params)) | |
| elif inst.gate == "ROTR": | |
| shift, width = int(inst.params[0]), int(inst.params[1]) | |
| inverted.append(QIRInstruction("ROTR", inst.controls, inst.targets, ((width - shift) % width, width))) | |
| else: | |
| inverted.append(QIRInstruction(f"{inst.gate}_DAGGER", inst.controls, inst.targets, inst.params)) | |
| return inverted | |
| def expand_unroll_directives(source: str) -> str: | |
| pattern = re.compile(r"#unroll\s+(\d+)\s+for\s+(\w+)\s+in\s+(\d+)\.\.(\d+)\s*\{", re.M) | |
| while True: | |
| match = pattern.search(source) | |
| if not match: | |
| return source | |
| count = int(match.group(1)) | |
| var = match.group(2) | |
| start = int(match.group(3)) | |
| end = int(match.group(4)) | |
| body_start = match.end() | |
| depth = 1 | |
| pos = body_start | |
| while pos < len(source) and depth: | |
| if source[pos] == "{": | |
| depth += 1 | |
| elif source[pos] == "}": | |
| depth -= 1 | |
| pos += 1 | |
| body = source[body_start : pos - 1] | |
| if count != end - start + 1: | |
| raise ValueError("unroll count must match inclusive range length") | |
| expanded = "\n".join( | |
| body.replace(f"{{{var}}}", str(value)).replace(f"${var}", str(value)) | |
| for value in range(start, end + 1) | |
| ) | |
| source = source[: match.start()] + expanded + source[pos:] | |
| def compile_source(source: str) -> List[QIRInstruction]: | |
| return QLambdaCompiler(Parser(Lexer(source).tokenize()).parse()).compile() | |