diff --git a/lib/Dialect/LWE/Conversions/LWEToLattigo/LWEToLattigo.cpp b/lib/Dialect/LWE/Conversions/LWEToLattigo/LWEToLattigo.cpp index c60fa76d9d..48540eee89 100644 --- a/lib/Dialect/LWE/Conversions/LWEToLattigo/LWEToLattigo.cpp +++ b/lib/Dialect/LWE/Conversions/LWEToLattigo/LWEToLattigo.cpp @@ -697,6 +697,97 @@ struct ConvertOrionChebyshevOp } }; +struct ConvertKernelLinearTransformOp + : public OpConversionPattern { + using OpConversionPattern::OpConversionPattern; + + LogicalResult matchAndRewrite( + kernel::LinearTransformOp op, OpAdaptor adaptor, + ConversionPatternRewriter& rewriter) const override { + LLVM_DEBUG(llvm::dbgs() << "Lowering Kernel LinearTransformOp\n"); + + FailureOr evaluatorResult = + getContextualEvaluator(op.getOperation()); + if (failed(evaluatorResult)) { + return rewriter.notifyMatchFailure( + op, "CKKS evaluator not found in function context"); + } + Value evaluator = evaluatorResult.value(); + + FailureOr encoderResult = + getContextualEvaluator(op.getOperation()); + if (failed(encoderResult)) { + return rewriter.notifyMatchFailure( + op, "CKKS encoder not found in function context"); + } + Value encoder = encoderResult.value(); + + // Extract level from input LWE ciphertext type + auto lweType = dyn_cast(op.getInput().getType()); + if (!lweType) { + return rewriter.notifyMatchFailure(op, "input is not LWE ciphertext"); + } + auto modulusChain = lweType.getModulusChain(); + if (!modulusChain) { + return rewriter.notifyMatchFailure(op, + "input LWE type has no modulus chain"); + } + int64_t levelQ = + modulusChain.getElements().size() - 1 - modulusChain.getCurrent(); + + // Convert diagonal_indices from I64 to I32 (Lattigo CKKSLinearTransformOp + // expects I32) + auto diagonalIndicesAttr = op.getDiagonalIndices(); + std::vector diagonalIndicesI32; + for (auto val : diagonalIndicesAttr) { + diagonalIndicesI32.push_back(static_cast(val)); + } + auto diagonalIndicesI32Attr = + rewriter.getDenseI32ArrayAttr(diagonalIndicesI32); + + // logBabyStepGiantStepRatio + // For now default to 0. + int64_t logBSGSRatio = 0; + + auto levelQAttr = rewriter.getI64IntegerAttr(levelQ); + auto logBSGSRatioAttr = rewriter.getI64IntegerAttr(logBSGSRatio); + + auto diagonalsAttr = op.getDiagonals(); + Value diagonalsValue = + rewriter.create(op.getLoc(), diagonalsAttr); + + auto linearTransformOp = rewriter.create( + op.getLoc(), adaptor.getInput().getType(), evaluator, encoder, + adaptor.getInput(), diagonalsValue, diagonalIndicesI32Attr, levelQAttr, + logBSGSRatioAttr); + + auto outputLweType = + dyn_cast(op.getResult().getType()); + if (!outputLweType) { + return rewriter.notifyMatchFailure(op, "output is not LWE ciphertext"); + } + auto outputModulusChain = outputLweType.getModulusChain(); + if (!outputModulusChain) { + return rewriter.notifyMatchFailure( + op, "output LWE type has no modulus chain"); + } + + Value result = linearTransformOp.getResult(); + if (outputModulusChain.getCurrent() < modulusChain.getCurrent()) { + int64_t diff = + modulusChain.getCurrent() - outputModulusChain.getCurrent(); + for (int64_t i = 0; i < diff; ++i) { + auto rescaleOp = rewriter.create( + op.getLoc(), result.getType(), evaluator, result); + result = rescaleOp.getResult(); + } + } + + rewriter.replaceOp(op, result); + return success(); + } +}; + struct ConvertKernelEvalChebyshevOp : public OpConversionPattern { using OpConversionPattern::OpConversionPattern; @@ -918,7 +1009,7 @@ struct LWEToLattigo : public impl::LWEToLattigoBase { .addIllegalOp(); + kernel::EvalChebyshevOp, kernel::LinearTransformOp>(); RewritePatternSet patterns(context); addStructuralConversionPatterns(typeConverter, patterns, target); @@ -1098,8 +1189,8 @@ struct LWEToLattigo : public impl::LWEToLattigoBase { ConvertCKKSEncryptOp, ConvertCKKSDecryptOp, ConvertCKKSEncodeOp, ConvertCKKSDecodeOp, ConvertCKKSLevelReduceOp, ConvertCKKSBootstrappingOp, ConvertOrionLinearTransformOp, - ConvertOrionChebyshevOp, ConvertKernelEvalChebyshevOp>(typeConverter, - context); + ConvertOrionChebyshevOp, ConvertKernelEvalChebyshevOp, + ConvertKernelLinearTransformOp>(typeConverter, context); } // Misc diff --git a/lib/Dialect/Secret/Conversions/SecretToBGV/SecretToBGV.cpp b/lib/Dialect/Secret/Conversions/SecretToBGV/SecretToBGV.cpp index 651dbe2474..f5eb337c67 100644 --- a/lib/Dialect/Secret/Conversions/SecretToBGV/SecretToBGV.cpp +++ b/lib/Dialect/Secret/Conversions/SecretToBGV/SecretToBGV.cpp @@ -100,7 +100,10 @@ class SecretToBGVTypeConverter ring(rlweRing), plaintextModulus(ptm), isBFV(isBFV) { - addConversion([](Type type, Attribute attr) { return type; }); + addConversion([](Type type, Attribute attr) -> std::optional { + if (isa(type)) return std::nullopt; + return type; + }); addConversion( [this](RankedTensorType type, mgmt::MgmtAttr mgmtAttr) -> Type { // For cases like tensor.empty + mgmt.init, we need to convert this @@ -217,22 +220,16 @@ struct SecretToBGV : public impl::SecretToBGVBase { bool usePublicKey = schemeParamAttr.getEncryptionType() == bgv::BGVEncryptionType::pk; - // NOTE: 2 ** logN != minSlotCount - // they have different semantic - // auto logN = schemeParamAttr.getLogN(); auto plaintextModulus = schemeParamAttr.getPlaintextModulus(); - - // pass option minSlotCount is actually the number of slots - // TODO(#1402): use a proper name for BGV auto rlweRing = getRlweRNSRing(context, schemeParamAttr.getQ().asArrayRef(), - minSlotCount); + 1 << schemeParamAttr.getLogN()); if (failed(rlweRing)) { return signalPassFailure(); } // Ensure that all secret types are uniform and have last dimension - // matching the ring parameter size. In other words, this asserts that any - // data-semantic tensors have been converted to ciphertext-semantic tensors - // with the correct shape. + // less than or equal to the ring parameter size. In other words, this + // asserts that any data-semantic tensors have been converted to + // ciphertext-semantic tensors with the correct shape. Operation* foundOp = walkAndDetect(module, [&](Operation* op) { ValueRange valuesToCheck = op->getOperands(); if (auto funcOp = dyn_cast(op)) { @@ -241,7 +238,7 @@ struct SecretToBGV : public impl::SecretToBGVBase { for (auto value : valuesToCheck) { if (auto secretTy = dyn_cast(value.getType())) { auto tensorTy = dyn_cast(secretTy.getValueType()); - if (tensorTy && tensorTy.getDimSize(tensorTy.getRank() - 1) != + if (tensorTy && tensorTy.getDimSize(tensorTy.getRank() - 1) > rlweRing.value() .getPolynomialModulus() .getPolynomial() @@ -255,7 +252,7 @@ struct SecretToBGV : public impl::SecretToBGVBase { if (foundOp != nullptr) { foundOp->emitError( "expected secret types to be tensors with last dimension " - "matching ring parameter"); + "less than or equal to ring parameter"); signalPassFailure(); return; } diff --git a/lib/Dialect/Secret/Conversions/SecretToCKKS/SecretToCKKS.cpp b/lib/Dialect/Secret/Conversions/SecretToCKKS/SecretToCKKS.cpp index ffbd8863b3..5ef29fd990 100644 --- a/lib/Dialect/Secret/Conversions/SecretToCKKS/SecretToCKKS.cpp +++ b/lib/Dialect/Secret/Conversions/SecretToCKKS/SecretToCKKS.cpp @@ -1,7 +1,9 @@ #include "lib/Dialect/Secret/Conversions/SecretToCKKS/SecretToCKKS.h" +#include #include #include +#include #include #include #include @@ -100,7 +102,10 @@ class SecretToCKKSTypeConverter SecretToCKKSTypeConverter(MLIRContext* ctx, polynomial::RingAttr rlweRing) : UniquelyNamedAttributeAwareTypeConverter( mgmt::MgmtDialect::kArgMgmtAttrName) { - addConversion([](Type type, Attribute attr) { return type; }); + addConversion([](Type type, Attribute attr) -> std::optional { + if (isa(type)) return std::nullopt; + return type; + }); addConversion( [this](RankedTensorType type, mgmt::MgmtAttr mgmtAttr) -> Type { // For cases like tensor.empty + mgmt.init, we need to convert this @@ -244,6 +249,118 @@ class SecretGenericPlaintextDivision } }; +struct LinearTransformOpConversion + : public ContextAwareOpConversionPattern { + LinearTransformOpConversion(const ContextAwareTypeConverter& typeConverter, + MLIRContext* context, int64_t ringDim, + PatternBenefit benefit = 1) + : ContextAwareOpConversionPattern(typeConverter, + context, benefit), + ringDim(ringDim) {} + + LogicalResult matchAndRewrite( + secret::GenericOp op, OpAdaptor adaptor, + ContextAwareConversionPatternRewriter& rewriter) const override { + if (op.getBody()->getOperations().size() > 2) { + return failure(); + } + + auto& innerOp = op.getBody()->getOperations().front(); + auto ltOp = dyn_cast(innerOp); + if (!ltOp) { + return failure(); + } + + // Convert inputs + SmallVector inputs; + for (Value operand : ltOp->getOperands()) { + if (auto* secretArg = op.getOpOperandForBlockArgument(operand)) { + inputs.push_back(adaptor.getInputs()[secretArg->getOperandNumber()]); + } else { + inputs.push_back(operand); + } + } + + // Convert result types + SmallVector resultTypes; + if (failed(getTypeConverter()->convertTypes(op.getResultTypes(), + op.getResults(), resultTypes))) + return failure(); + + // Preserve attributes (similar to SecretGenericOpConversion) + SmallVector attrsToPreserve; + for (auto& namedAttr : ltOp->getDialectAttrs()) { + attrsToPreserve.push_back(namedAttr); + } + for (auto attrName : ltOp.getAttributeNames()) { + if (attrName == "diagonals") + continue; // We will handle diagonals separately + if (auto attr = ltOp->getAttr(attrName)) { + attrsToPreserve.push_back(rewriter.getNamedAttr(attrName, attr)); + } + } + + // Pad diagonals + auto diagonalsAttr = cast(ltOp.getDiagonals()); + auto diagonalsType = cast(diagonalsAttr.getType()); + auto shape = diagonalsType.getShape(); + int64_t numDiagonals = shape[0]; + int64_t numCols = shape[1]; + + int64_t actualSlots = ringDim / 2; // CKKS assumption + + DenseElementsAttr newDiagonalsAttr; + if (numCols == actualSlots) { + newDiagonalsAttr = diagonalsAttr; + } else { + if (numCols > actualSlots) { + return ltOp.emitOpError("diagonals slot size (") + << numCols << ") is larger than actual slots (" << actualSlots + << ")"; + } + SmallVector paddedValues; + auto elementValues = diagonalsAttr.getValues(); + auto elemType = diagonalsType.getElementType(); + Attribute zeroAttr = rewriter.getZeroAttr(elemType); + + for (int64_t i = 0; i < numDiagonals; ++i) { + for (int64_t j = 0; j < numCols; ++j) { + paddedValues.push_back(elementValues[i * numCols + j]); + } + for (int64_t j = numCols; j < actualSlots; ++j) { + paddedValues.push_back(zeroAttr); + } + } + + auto newDiagonalsType = + RankedTensorType::get({numDiagonals, actualSlots}, elemType); + newDiagonalsAttr = DenseElementsAttr::get(newDiagonalsType, paddedValues); + } + attrsToPreserve.push_back( + rewriter.getNamedAttr("diagonals", newDiagonalsAttr)); + + // Handle mgmt attrs + convertArrayOfDicts(op.getAllResultAttrsAttr(), attrsToPreserve); + convertArrayOfDicts(op.getAllOperandAttrsAttr(), attrsToPreserve); + DenseSet seenNames; + SmallVector dedupedAttrsToPreserve; + for (auto attr : llvm::reverse(attrsToPreserve)) { + if (seenNames.insert(attr.getName().getValue()).second) { + dedupedAttrsToPreserve.push_back(attr); + } + } + std::reverse(dedupedAttrsToPreserve.begin(), dedupedAttrsToPreserve.end()); + auto newLtOp = kernel::LinearTransformOp::create( + rewriter, ltOp.getLoc(), resultTypes, inputs, dedupedAttrsToPreserve); + + rewriter.replaceOp(op, newLtOp->getResults()); + return success(); + } + + private: + int64_t ringDim; +}; + struct SecretToCKKS : public impl::SecretToCKKSBase { using SecretToCKKSBase::SecretToCKKSBase; @@ -266,7 +383,7 @@ struct SecretToCKKS : public impl::SecretToCKKSBase { // pass option minSlotCount is actually the number of slots // TODO(#1402): use a proper name for CKKS auto rlweRing = getRlweRNSRing(context, schemeParamAttr.getQ().asArrayRef(), - minSlotCount); + 1 << schemeParamAttr.getLogN()); if (failed(rlweRing)) { return signalPassFailure(); } @@ -310,6 +427,9 @@ struct SecretToCKKS : public impl::SecretToCKKSBase { SecretGenericOpLevelReduceConversion>( typeConverter, context); + int64_t ringDim = 1 << schemeParamAttr.getLogN(); + patterns.add(typeConverter, context, ringDim); + patterns.add(typeConverter, context, usePublicKey, rlweRing.value()); patterns.add(typeConverter, context, rlweRing.value()); diff --git a/lib/Target/Lattigo/LattigoEmitter.cpp b/lib/Target/Lattigo/LattigoEmitter.cpp index 2d80b6db57..33ba6069e7 100644 --- a/lib/Target/Lattigo/LattigoEmitter.cpp +++ b/lib/Target/Lattigo/LattigoEmitter.cpp @@ -2080,7 +2080,7 @@ LogicalResult LattigoEmitter::printOperation(CKKSEncodeOp op) { auto valueName = getName(op.getValue()); auto maxSlotsName = getName(newPlaintextOp.getParams()) + ".MaxSlots()"; auto numSlotsAttr = dyn_cast_or_null( - op->getParentOfType()->getAttr(kRequestedSlotCountAttrName)); + op->getParentOfType()->getAttr(kActualSlotCountAttrName)); if (numSlotsAttr) { maxSlotsName = std::to_string(numSlotsAttr.getInt()); imports.insert(std::string(kRingImport)); @@ -2302,6 +2302,13 @@ LogicalResult LattigoEmitter::printOperation(CKKSLinearTransformOp op) { } int64_t slotsPerDiagonal = diagonalsType.getShape()[1]; + Type elementType = diagonalsType.getElementType(); + bool isF64 = false; + if (auto floatType = dyn_cast(elementType)) { + if (floatType.getWidth() == 64) { + isF64 = true; + } + } // Generate unique variable names std::string diagonalsMapName = outputName + "_diags"; @@ -2310,14 +2317,28 @@ LogicalResult LattigoEmitter::printOperation(CKKSLinearTransformOp op) { std::string ltName = outputName + "_lt"; std::string ltEvalName = outputName + "_lteval"; std::string errName = getErrName(); + std::string slotsName = outputName + "_slots"; os << diagonalIndices << " := " << printDenseI32ArrayAttr(op.getDiagonalIndicesAttr()) << "\n"; + os << slotsName << " := 1 << " << inputName << ".LogDimensions.Cols\n"; os << diagonalsMapName << " := make(lintrans.Diagonals[float64])\n"; os << "for i, diagIndex := range " << diagonalIndices << " {\n"; os.indent(); - os << diagonalsMapName << "[diagIndex] = " << diagonalsName << "[i*" - << slotsPerDiagonal << ":(i+1)*" << slotsPerDiagonal << "]\n"; + if (isF64) { + os << diagonalsMapName << "[diagIndex] = " << diagonalsName << "[i*" + << slotsPerDiagonal << ":i*" << slotsPerDiagonal << " + " << slotsName + << "]\n"; + } else { + os << "diag := make([]float64, " << slotsName << ")\n"; + os << "for j := 0; j < " << slotsName << "; j++ {\n"; + os.indent(); + os << "diag[j] = float64(" << diagonalsName << "[i*" << slotsPerDiagonal + << " + j])\n"; + os.unindent(); + os << "}\n"; + os << diagonalsMapName << "[diagIndex] = diag\n"; + } os.unindent(); os << "}\n"; @@ -2325,10 +2346,10 @@ LogicalResult LattigoEmitter::printOperation(CKKSLinearTransformOp op) { os.indent(); os << "DiagonalsIndexList: " << diagonalsMapName << ".DiagonalsIndexList(),\n"; - os << "LevelQ: " << op.getLevelQ().getInt() << ",\n"; + os << "LevelQ: " << inputName << ".Level(),\n"; os << "LevelP: " << evaluatorName << ".GetRLWEParameters().MaxLevelP(),\n"; os << "Scale: rlwe.NewScale(" << evaluatorName << ".GetRLWEParameters().Q()[" - << op.getLevelQ().getInt() << "]),\n"; + << inputName << ".Level()]),\n"; os << "LogDimensions: " << inputName << ".LogDimensions,\n"; os << "LogBabyStepGiantStepRatio: " << op.getLogBabyStepGiantStepRatio().getInt() << ",\n"; diff --git a/lib/Target/Lattigo/TargetConfig.td b/lib/Target/Lattigo/TargetConfig.td index c77e32e4fc..27469aab4e 100644 --- a/lib/Target/Lattigo/TargetConfig.td +++ b/lib/Target/Lattigo/TargetConfig.td @@ -3,4 +3,5 @@ include "lib/Target/CompilationTarget/HEIRTarget.td" def Lattigo : CompilationTarget<"lattigo"> { let bootstrapLevelsConsumed = 16; let has_kernel_chebyshev = 1; + let has_kernel_linear_transform = 1; } diff --git a/lib/Transforms/ConvertToCiphertextSemantics/ConvertToCiphertextSemantics.cpp b/lib/Transforms/ConvertToCiphertextSemantics/ConvertToCiphertextSemantics.cpp index 5fb74d9260..3b285fa4ea 100644 --- a/lib/Transforms/ConvertToCiphertextSemantics/ConvertToCiphertextSemantics.cpp +++ b/lib/Transforms/ConvertToCiphertextSemantics/ConvertToCiphertextSemantics.cpp @@ -966,6 +966,12 @@ struct PreserveLinalgMatvecAsLinearTransform } int64_t numDiagonals = convertedMatrixType.getShape()[0]; + int64_t numCols = matrixType.getDimSize(1); + if (numDiagonals < numCols) { + return rewriter.notifyMatchFailure( + op, "requires post-processing (numDiagonals < numCols)"); + } + int64_t slots = convertedMatrixType.getShape()[1]; auto elementType = matrixType.getElementType(); @@ -975,8 +981,6 @@ struct PreserveLinalgMatvecAsLinearTransform auto matrixRelation = matrixLayout.getIntegerRelation(); PointPairCollector collector(2, 2); enumeratePoints(matrixRelation, collector); - - int64_t numCols = matrixType.getDimSize(1); for (const auto& pointPair : collector.points) { int64_t row = pointPair.first[0]; int64_t col = pointPair.first[1]; diff --git a/tests/Dialect/LWE/Conversions/lwe_to_lattigo/linear_transform.mlir b/tests/Dialect/LWE/Conversions/lwe_to_lattigo/linear_transform.mlir new file mode 100644 index 0000000000..a4853a997a --- /dev/null +++ b/tests/Dialect/LWE/Conversions/lwe_to_lattigo/linear_transform.mlir @@ -0,0 +1,25 @@ +// RUN: heir-opt --lwe-to-lattigo %s | FileCheck %s + +#inverse_canonical_encoding = #lwe.inverse_canonical_encoding +#key = #lwe.key<> +#modulus_chain = #lwe.modulus_chain, current = 0> +#ring_f64_1_x1024 = #polynomial.ring> +!rns_L0 = !rns.rns> +#ring_rns_L0_1_x1024 = #polynomial.ring> +#ciphertext_space_L0 = #lwe.ciphertext_space +!ct = !lwe.lwe_ciphertext, ciphertext_space = #ciphertext_space_L0, key = #key, modulus_chain = #modulus_chain> + +// CHECK: ![[CT:.*]] = !lattigo.rlwe.ciphertext +// CHECK: ![[ENCODER:.*]] = !lattigo.ckks.encoder +// CHECK: ![[EVAL:.*]] = !lattigo.ckks.evaluator + +module attributes {backend.lattigo, ckks.schemeParam = #ckks.scheme_param, scheme.ckks} { + // CHECK: func.func @test_linear_transform(%[[EVAL:.*]]: ![[EVAL]], %{{.*}}: {{.*}}, %[[ENCODER:.*]]: ![[ENCODER]], %[[CT:.*]]: ![[CT]]) -> ![[CT]] + // CHECK: %[[DIAGONALS:.*]] = arith.constant dense<{{\[\[}}1.000000e+00, 2.000000e+00], [3.000000e+00, 4.000000e+00{{\]\]}}> : tensor<2x2xf64> + // CHECK: %[[VAL_3:.*]] = lattigo.ckks.linear_transform %[[EVAL]], %[[ENCODER]], %[[CT]], %[[DIAGONALS]] {diagonal_indices = array, levelQ = 1 : i64, logBabyStepGiantStepRatio = 0 : i64} : (![[EVAL]], ![[ENCODER]], ![[CT]], tensor<2x2xf64>) -> ![[CT]] + // CHECK: return %[[VAL_3]] : ![[CT]] + func.func @test_linear_transform(%ct: !ct) -> !ct { + %0 = kernel.linear_transform %ct {diagonals = dense<[[1.0, 2.0], [3.0, 4.0]]> : tensor<2x2xf64>, diagonal_indices = array} : !ct -> !ct + return %0 : !ct + } +} diff --git a/tests/Dialect/Secret/Conversions/secret_to_bgv/invalid.mlir b/tests/Dialect/Secret/Conversions/secret_to_bgv/invalid.mlir index 7b19c5870f..d865ff2c39 100644 --- a/tests/Dialect/Secret/Conversions/secret_to_bgv/invalid.mlir +++ b/tests/Dialect/Secret/Conversions/secret_to_bgv/invalid.mlir @@ -4,9 +4,9 @@ #mgmt = #mgmt.mgmt module attributes {bgv.schemeParam = #bgv.scheme_param} { - // expected-error@below {{expected secret types to be tensors with last dimension matching ring parameter}} - func.func @test_invalid_dimension(%arg0 : !secret.secret> {mgmt.mgmt = #mgmt}) -> (!secret.secret> {mgmt.mgmt = #mgmt}) { - return %arg0 : !secret.secret> + // expected-error@below {{expected secret types to be tensors with last dimension less than or equal to ring parameter}} + func.func @test_invalid_dimension(%arg0 : !secret.secret> {mgmt.mgmt = #mgmt}) -> (!secret.secret> {mgmt.mgmt = #mgmt}) { + return %arg0 : !secret.secret> } } diff --git a/tests/Dialect/Secret/Conversions/secret_to_bgv/ops.mlir b/tests/Dialect/Secret/Conversions/secret_to_bgv/ops.mlir index 503186ebd8..d0255bec0a 100644 --- a/tests/Dialect/Secret/Conversions/secret_to_bgv/ops.mlir +++ b/tests/Dialect/Secret/Conversions/secret_to_bgv/ops.mlir @@ -20,7 +20,7 @@ module attributes {bgv.schemeParam = #bgv.scheme_param } -> (!eui1 {mgmt.mgmt = #mgmt1}) // CHECK: return - // CHECK-SAME: polynomialModulus = <1 + x**1024> + // CHECK-SAME: polynomialModulus = <1 + x**16384> // CHECK-SAME: size = 3 return %1 : !eui1 } diff --git a/tests/Dialect/Secret/Conversions/secret_to_ckks/ops.mlir b/tests/Dialect/Secret/Conversions/secret_to_ckks/ops.mlir index 854cdf9fb7..08e58277d1 100644 --- a/tests/Dialect/Secret/Conversions/secret_to_ckks/ops.mlir +++ b/tests/Dialect/Secret/Conversions/secret_to_ckks/ops.mlir @@ -22,7 +22,7 @@ module attributes {ckks.schemeParam = #ckks.scheme_param } -> (!eui1 {mgmt.mgmt = #mgmt1}) // CHECK: return - // CHECK-SAME: polynomialModulus = <1 + x**1024> + // CHECK-SAME: polynomialModulus = <1 + x**16384> // CHECK-SAME: size = 3 return %1 : !eui1 } @@ -48,7 +48,7 @@ module attributes {ckks.schemeParam = #ckks.scheme_param } -> (!efi1 {mgmt.mgmt = #mgmt1}) // CHECK: return - // CHECK-SAME: polynomialModulus = <1 + x**1024> + // CHECK-SAME: polynomialModulus = <1 + x**16384> // CHECK-SAME: size = 3 return %3 : !efi1 } @@ -64,7 +64,7 @@ module attributes {ckks.schemeParam = #ckks.scheme_param } -> (!secret.secret> {mgmt.mgmt = #mgmt}) // CHECK: return - // CHECK-SAME: coefficientType = !rns.rns>, polynomialModulus = <1 + x**1024> + // CHECK-SAME: coefficientType = !rns.rns>, polynomialModulus = <1 + x**16384> return %0 : !secret.secret> } } diff --git a/tests/Examples/lattigo/ckks/matvec_512x784/matvec_512x784_test.go b/tests/Examples/lattigo/ckks/matvec_512x784/matvec_512x784_test.go index 97a4c848f4..1aa54c2aa5 100644 --- a/tests/Examples/lattigo/ckks/matvec_512x784/matvec_512x784_test.go +++ b/tests/Examples/lattigo/ckks/matvec_512x784/matvec_512x784_test.go @@ -17,10 +17,12 @@ func TestMatvec(t *testing.T) { expected := float32(78.4) ct0 := Matvec__encrypt__arg0(evaluator, params, ecd, enc, arg0) - zero := Matvec__encrypt__zero__0(evaluator, params, ecd, enc) - resultCt := Matvec(evaluator, params, ecd, ct0, zero) + zeroCt := Matvec__encrypt__zero__0(evaluator, params, ecd, enc) + resultCt := Matvec(evaluator, params, ecd, ct0, zeroCt) result := Matvec__decrypt__result0(evaluator, params, ecd, dec, resultCt) - errorThreshold := float64(2.0) + // Error threshold increased to 4.0 due to fallback to Halevi-Shoup kernel + // which has different noise characteristics. + errorThreshold := float64(4.0) for i := 0; i < rows; i++ { if math.Abs(float64(result[i]-expected)) > errorThreshold { t.Errorf("Decryption error at index %d: %.2f != %.2f", i, result[i], expected) diff --git a/tests/Examples/lattigo/ckks/matvec_square/matvec_512x512_test.go b/tests/Examples/lattigo/ckks/matvec_square/matvec_512x512_test.go index c16c303be0..f7b4e53cc1 100644 --- a/tests/Examples/lattigo/ckks/matvec_square/matvec_512x512_test.go +++ b/tests/Examples/lattigo/ckks/matvec_square/matvec_512x512_test.go @@ -17,8 +17,7 @@ func TestMatvec(t *testing.T) { expected := float32(51.2) ct0 := Matvec__encrypt__arg0(evaluator, params, ecd, enc, arg0) - zero := Matvec__encrypt__zero__0(evaluator, params, ecd, enc) - resultCt := Matvec(evaluator, params, ecd, ct0, zero) + resultCt := Matvec(evaluator, params, ecd, ct0) result := Matvec__decrypt__result0(evaluator, params, ecd, dec, resultCt) errorThreshold := float64(0.00001) for i := 0; i < rows; i++ { diff --git a/tests/Examples/openfhe/ckks/halevi_shoup_matvec/BUILD b/tests/Examples/openfhe/ckks/halevi_shoup_matvec/BUILD index 8cde5ffde3..787b753d81 100644 --- a/tests/Examples/openfhe/ckks/halevi_shoup_matvec/BUILD +++ b/tests/Examples/openfhe/ckks/halevi_shoup_matvec/BUILD @@ -7,7 +7,7 @@ openfhe_end_to_end_test( generated_lib_header = "halevi_shoup_matvec_lib.h", heir_opt_flags = [ "--annotate-module=backend=openfhe scheme=ckks", - "--torch-linalg-to-ckks=min-slot-count=8192", + "--torch-linalg-to-ckks=min-slot-count=1024", "--scheme-to-openfhe", ], heir_translate_flags = [ diff --git a/tests/Transforms/convert_to_ciphertext_semantics/linear_transform.mlir b/tests/Transforms/convert_to_ciphertext_semantics/linear_transform.mlir index d62da35cc6..a4c96d0a1c 100644 --- a/tests/Transforms/convert_to_ciphertext_semantics/linear_transform.mlir +++ b/tests/Transforms/convert_to_ciphertext_semantics/linear_transform.mlir @@ -1,4 +1,4 @@ -// RUN: heir-opt %s --convert-to-ciphertext-semantics=min-slot-count=4 | FileCheck %s +// RUN: heir-opt %s --split-input-file --convert-to-ciphertext-semantics=min-slot-count=8 | FileCheck %s // CHECK: module module attributes { @@ -21,7 +21,7 @@ module attributes { ^body(%input: tensor<4xf32>): // CHECK: kernel.linear_transform // CHECK-SAME: diagonal_indices = array - // CHECK-SAME: diagonals = dense<{{\[\[}}1.000000e+00, 6.000000e+00, 0.000000e+00, 0.000000e+00], [2.000000e+00, 7.000000e+00, 0.000000e+00, 0.000000e+00], [3.000000e+00, 8.000000e+00, 0.000000e+00, 0.000000e+00], [4.000000e+00, 5.000000e+00, 0.000000e+00, 0.000000e+00]]> : tensor<4x4xf32> + // CHECK-SAME: diagonals = dense<{{\[\[}}1.000000e+00, 6.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00], [2.000000e+00, 7.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00], [3.000000e+00, 8.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00], [4.000000e+00, 5.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00, 0.000000e+00]]> : tensor<4x8xf32> %3 = linalg.matvec { secret.kernel = #secret.kernel, tensor_ext.layout = #tensor_ext.layout<"{ [i0] -> [ct, slot] : ct = 0 and slot = i0 and 0 <= i0 <= 1 and 0 <= slot <= 3 }"> @@ -31,3 +31,34 @@ module attributes { return %2 : !secret.secret> } } + +// ----- + +#kernel = #secret.kernel +#layout = #tensor_ext.layout<"{ [i0] -> [ct, slot] : ct = 0 and (-i0 + slot) mod 4 = 0 and 0 <= i0 <= 3 and 0 <= slot <= 7 }"> +#layout1 = #tensor_ext.layout<"{ [i0, i1] -> [ct, slot] : (i0 - i1 + ct) mod 4 = 0 and (-i0 + slot) mod 4 = 0 and 0 <= i0 <= 3 and 0 <= i1 <= 3 and 0 <= ct <= 3 and 0 <= slot <= 7 }"> + +module attributes { + backend.openfhe, + backend.config_override = {has_kernel_linear_transform = true} +} { + // CHECK: func.func @matvec_to_linear_transform + // CHECK-SAME: (%[[ARG0:.*]]: !secret.secret> {{.*}}) -> (!secret.secret> {{.*}}) + // CHECK: secret.generic + // CHECK: kernel.linear_transform {{%[a-zA-Z0-9_]+}} + // CHECK-SAME: diagonal_indices = array + // CHECK-SAME: diagonals = dense<{{.*}}> : tensor<4x8xf32> + // CHECK: secret.yield + func.func @matvec_to_linear_transform(%arg0: !secret.secret> {tensor_ext.layout = #layout}) -> (!secret.secret> {tensor_ext.layout = #layout}) { + %cst = arith.constant dense<0.000000e+00> : tensor<4xf32> + %cst_0 = arith.constant dense<[[1.0, 2.0, 3.0, 4.0], [5.0, 1.0, 2.0, 3.0], [6.0, 5.0, 1.0, 2.0], [7.0, 6.0, 5.0, 1.0]]> : tensor<4x4xf32> + %0 = secret.generic(%arg0: !secret.secret> {tensor_ext.layout = #layout}) { + ^body(%input0: tensor<4xf32>): + %1 = tensor_ext.assign_layout %cst_0 {layout = #layout1, tensor_ext.layout = #layout1} : tensor<4x4xf32> + %2 = tensor_ext.assign_layout %cst {layout = #layout, tensor_ext.layout = #layout} : tensor<4xf32> + %3 = linalg.matvec {secret.kernel = #kernel, tensor_ext.layout = #layout} ins(%1, %input0 : tensor<4x4xf32>, tensor<4xf32>) outs(%2 : tensor<4xf32>) -> tensor<4xf32> + secret.yield %3 : tensor<4xf32> + } -> (!secret.secret> {tensor_ext.layout = #layout}) + return %0 : !secret.secret> + } +} diff --git a/tests/Transforms/convert_to_ciphertext_semantics/matvec_512x784.mlir b/tests/Transforms/convert_to_ciphertext_semantics/matvec_512x784.mlir index 864d3ba3dd..8f479f4583 100644 --- a/tests/Transforms/convert_to_ciphertext_semantics/matvec_512x784.mlir +++ b/tests/Transforms/convert_to_ciphertext_semantics/matvec_512x784.mlir @@ -23,7 +23,11 @@ #layout = #tensor_ext.layout<"{ [i0] -> [ct, slot] : ct = 0 and (-i0 + slot) mod 512 = 0 and 0 <= i0 <= 511 and 0 <= slot <= 1023 }"> #layout1 = #tensor_ext.layout<"{ [i0] -> [ct, slot] : ct = 0 and (-i0 + slot) mod 1024 = 0 and 0 <= i0 <= 783 and 0 <= slot <= 1023 }"> #layout2 = #tensor_ext.layout<"{ [i0, i1] -> [ct, slot] : (i0 - i1 + ct) mod 512 = 0 and (-i1 + ct + slot) mod 1024 = 0 and 0 <= i0 <= 511 and 0 <= i1 <= 783 and 0 <= ct <= 511 and 0 <= slot <= 1023 }"> -module attributes {backend.lattigo, scheme.ckks} { +module attributes { + backend.lattigo, + backend.config_override = {has_kernel_linear_transform = false}, + scheme.ckks +} { func.func @matvec(%arg0: !secret.secret> {tensor_ext.layout = #layout1}) -> (!secret.secret> {tensor_ext.layout = #layout}) { %cst = arith.constant dense<0.000000e+00> : tensor<512xf32> %cst_0 = arith.constant dense<1.000000e+00> : tensor<512x784xf32>