Skip to content

Commit 198c5fa

Browse files
committed
scratchalloc and moculecreate ops
1 parent e85648a commit 198c5fa

3 files changed

Lines changed: 32 additions & 4 deletions

File tree

lib/Target/Poulpy/PoulpyEmitter.cpp

Lines changed: 19 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -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+
574588
LogicalResult 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();

lib/Target/Poulpy/PoulpyEmitter.h

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -68,6 +68,8 @@ class PoulpyEmitter {
6868
LogicalResult printOperation(func::ReturnOp op);
6969
LogicalResult printOperation(func::CallOp op);
7070
LogicalResult printOperation(memref::AllocOp op);
71+
LogicalResult printOperation(ModuleCreateOp op);
72+
LogicalResult printOperation(ScratchAllocOp op);
7173
LogicalResult printOperation(AddOp op);
7274
LogicalResult printOperation(AddAssignOp op);
7375
LogicalResult printOperation(SubOp op);

tests/Emitter/Poulpy/emit_poulpy.mlir

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -502,3 +502,14 @@ func.func @call_two_results(%mod: !module, %s: !scratch, %a: !ct, %b: !ct) -> (!
502502
// CHECK-NEXT: Ok(([[r0]], [[r1]]))
503503
return %r0, %r1 : !ct, !ct
504504
}
505+
506+
// CHECK: pub fn setup(
507+
// CHECK-NEXT: ) -> Result<(Module<BE>, ScratchOwned<BE>)> {
508+
func.func @setup() -> (!module, !scratch) {
509+
// CHECK: let [[m:v[0-9]+]] = Module::<BE>::new(64u64);
510+
%mod = poulpy.module_create {N = 64 : i64} : () -> !module
511+
// CHECK-NEXT: let mut [[s:v[0-9]+]] = ScratchOwned::<BE>::alloc(1024usize);
512+
%scratch = poulpy.scratch_alloc {size = 1024 : i64} : () -> !scratch
513+
// CHECK-NEXT: Ok(([[m]], [[s]]))
514+
return %mod, %scratch : !module, !scratch
515+
}

0 commit comments

Comments
 (0)