@@ -99,8 +99,8 @@ LogicalResult PoulpyEmitter::translate(Operation& op) {
9999 AddAssignOp, SubOp, SubAssignOp, MulOp, MulAssignOp, RotateOp,
100100 RotateAssignOp, RescaleOp, RescaleAssignOp, CompactLimbsOp,
101101 AddUnnormalizedOp, SubUnnormalizedOp, NormalizeOp, EncodeOp,
102- DecodeOp, EncryptOp, DecryptOp, memref::AllocOp>(
103- [&](auto op) { return printOperation (op); })
102+ DecodeOp, EncryptOp, DecryptOp, ModuleCreateOp, ScratchAllocOp,
103+ memref::AllocOp>( [&](auto op) { return printOperation (op); })
104104 .Default ([&](Operation& op) {
105105 return op.emitOpError (" unable to find printer for op" );
106106 });
@@ -571,6 +571,20 @@ LogicalResult PoulpyEmitter::printOperation(DecryptOp decryptOp) {
571571 return success ();
572572}
573573
574+ LogicalResult PoulpyEmitter::printOperation (ModuleCreateOp moduleCreateOp) {
575+ os << " let " << variableNames->getNameForValue (moduleCreateOp.getModule ())
576+ << " = Module::<BE>::new(" << moduleCreateOp.getN () << " u64);\n " ;
577+ return success ();
578+ }
579+
580+ LogicalResult PoulpyEmitter::printOperation (ScratchAllocOp scratchAllocOp) {
581+ os << " let mut "
582+ << variableNames->getNameForValue (scratchAllocOp.getScratch ())
583+ << " = ScratchOwned::<BE>::alloc(" << scratchAllocOp.getSize ()
584+ << " usize);\n " ;
585+ return success ();
586+ }
587+
574588LogicalResult PoulpyEmitter::printOperation (memref::AllocOp allocOp) {
575589 MemRefType resultType = allocOp.getType ();
576590 if (resultType.getElementType ().isF64 ()) {
@@ -592,10 +606,11 @@ FailureOr<std::string> PoulpyEmitter::convertType(Type type, bool isArg,
592606 bool isMutated) {
593607 return llvm::TypeSwitch<Type&, FailureOr<std::string>>(type)
594608 .Case <ModuleType>([&](ModuleType) -> FailureOr<std::string> {
595- return std::string (" & Module<BE>" );
609+ return std::string (isArg ? " &Module<BE> " : " Module<BE>" );
596610 })
597611 .Case <ScratchType>([&](ScratchType) -> FailureOr<std::string> {
598- return std::string (" &mut ScratchOwned<BE>" );
612+ return std::string (isArg ? " &mut ScratchOwned<BE>"
613+ : " ScratchOwned<BE>" );
599614 })
600615 .Case <MemRefType>([&](MemRefType memRefType) -> FailureOr<std::string> {
601616 Type elementType = memRefType.getElementType ();
0 commit comments