diff --git a/cudaq/lib/Optimizer/Dialect/Quake/CanonicalPatterns.inc b/cudaq/lib/Optimizer/Dialect/Quake/CanonicalPatterns.inc index ce612b18be7..5cc1753fa7e 100644 --- a/cudaq/lib/Optimizer/Dialect/Quake/CanonicalPatterns.inc +++ b/cudaq/lib/Optimizer/Dialect/Quake/CanonicalPatterns.inc @@ -132,6 +132,23 @@ struct ForwardEmptyVeqSizePattern } }; +// %0 = quake.alloca !quake.veq<0> +// ───────────────────────────────── +// %0 = cc.undef !quake.veq<0> +struct ReplaceZeroSizeAllocaPattern + : public OpRewritePattern { + using OpRewritePattern::OpRewritePattern; + + LogicalResult matchAndRewrite(cudaq::quake::AllocaOp alloc, + PatternRewriter &rewriter) const override { + auto veqTy = dyn_cast(alloc.getType()); + if (!veqTy || !veqTy.hasSpecifiedSize() || veqTy.getSize() != 0) + return failure(); + rewriter.replaceOpWithNewOp(alloc, veqTy); + return success(); + } +}; + // %2 = constant 10 : i32 // %3 = quake.alloca !quake.veq[%2 : i32] // ───────────────────────────────────────── diff --git a/cudaq/lib/Optimizer/Dialect/Quake/QuakeOps.cpp b/cudaq/lib/Optimizer/Dialect/Quake/QuakeOps.cpp index f1d1a162649..112ce2201e6 100644 --- a/cudaq/lib/Optimizer/Dialect/Quake/QuakeOps.cpp +++ b/cudaq/lib/Optimizer/Dialect/Quake/QuakeOps.cpp @@ -319,7 +319,8 @@ void cudaq::quake::AllocaOp::getCanonicalizationPatterns( // Use a canonicalization pattern as folding the constant into the veq type // changes the type. Uses may still expect a veq with unspecified size. // Folding is strictly reductive and doesn't allow the creation of ops. - patterns.add(context); + patterns.add( + context); } cudaq::quake::InitializeStateOp cudaq::quake::AllocaOp::getInitializedState() { diff --git a/python/cudaq/kernel/ast_bridge.py b/python/cudaq/kernel/ast_bridge.py index 3ce6646e070..915aab4298f 100644 --- a/python/cudaq/kernel/ast_bridge.py +++ b/python/cudaq/kernel/ast_bridge.py @@ -4683,25 +4683,22 @@ def get_item_type(pyval): veqTy = self.getVeqType() c0 = self.getConstantInt(0) c1 = self.getConstantInt(1) - empty_veq_ty = quake.VeqType.get(0, context=self.ctx) - init_veq = quake.RelaxSizeOp( - veqTy, - quake.AllocaOp(empty_veq_ty).result).result + veq1_ty = quake.VeqType.get(1, context=self.ctx) + + def extractElem(i): + if quake.VeqType.isinstance(iterable.type): + return quake.ExtractRefOp(iterTy, iterable, -1, + index=i).result + elem_addr = cc.ComputePtrOp( + cc.PointerType.get(iterTy), iterable, [i], + DenseI32ArrayAttr.get([kDynamicPtrIndex], + context=self.ctx)) + return cc.LoadOp(elem_addr).result def bodyBuilder(args): i, curr_veq = args[0], args[1] - if quake.VeqType.isinstance(iterable.type): - idx_val = quake.ExtractRefOp(iterTy, - iterable, - -1, - index=i).result - else: - elem_addr = cc.ComputePtrOp( - cc.PointerType.get(iterTy), iterable, [i], - DenseI32ArrayAttr.get([kDynamicPtrIndex], - context=self.ctx)) - idx_val = cc.LoadOp(elem_addr).result + idx_val = extractElem(i) self.symbolTable.beginBlock() self.__deconstructAssignment(node.generators[0].target, idx_val) @@ -4726,12 +4723,60 @@ def bodyBuilder(args): self.symbolTable.endBlock() cc.ContinueOp([i, new_veq]) - loop = self.createForLoop( - [i64Ty, veqTy], bodyBuilder, [c0, init_veq], - lambda args: arith.CmpIOp(IntegerAttr.get(i64Ty, 2), args[ - 0], iterableSize).result, - lambda args: [arith.AddIOp(args[0], c1).result, args[1]]) - self.pushValue(loop.results[1]) + # Check if the source iterable is empty. If so, return an + # undefined veq<0>. Otherwise peel the first element to form a + # veq<1> seed so the loop never starts from a veq<0> alloca + # (which has no valid semantics). + isEmpty = arith.CmpIOp(IntegerAttr.get(i64Ty, 0), iterableSize, + c0).result + ifEmptyOp = cc.IfOp([veqTy], isEmpty, []) + emptyBlock = Block.create_at_start(ifEmptyOp.thenRegion, []) + with InsertionPoint(emptyBlock): + undef = cc.UndefOp(empty_veq_ty) + cc.ContinueOp( + [quake.RelaxSizeOp(veqTy, undef.result).result]) + + nonEmptyBlock = Block.create_at_start(ifEmptyOp.elseRegion, []) + with InsertionPoint(nonEmptyBlock): + # Peel element 0 to form the initial veq<1> seed. + idx_val_0 = extractElem(c0) + self.symbolTable.beginBlock() + self.__deconstructAssignment(node.generators[0].target, + idx_val_0) + if hasFilter: + cond0 = evalFilter() + ifFilter0 = cc.IfOp([veqTy], cond0, []) + thenBlock0 = Block.create_at_start( + ifFilter0.thenRegion, []) + with InsertionPoint(thenBlock0): + self.visit(node.elt) + ref0 = self.popValue() + veq1 = quake.ConcatOp(veq1_ty, [ref0]).result + cc.ContinueOp( + [quake.RelaxSizeOp(veqTy, veq1).result]) + elseBlock0 = Block.create_at_start( + ifFilter0.elseRegion, []) + with InsertionPoint(elseBlock0): + undef0 = cc.UndefOp(empty_veq_ty) + cc.ContinueOp([ + quake.RelaxSizeOp(veqTy, undef0.result).result + ]) + init_seed = ifFilter0.result + else: + self.visit(node.elt) + ref0 = self.popValue() + veq1 = quake.ConcatOp(veq1_ty, [ref0]).result + init_seed = quake.RelaxSizeOp(veqTy, veq1).result + self.symbolTable.endBlock() + + loop = self.createForLoop( + [i64Ty, veqTy], bodyBuilder, [c1, init_seed], + lambda args: arith.CmpIOp(IntegerAttr.get( + i64Ty, 2), args[0], iterableSize).result, lambda + args: [arith.AddIOp(args[0], c1).result, args[1]]) + cc.ContinueOp([loop.results[1]]) + + self.pushValue(ifEmptyOp.result) return self.emitFatalError( "unsupported list comprehension producing qubit references",