diff --git a/lib/Analysis/SecretnessAnalysis/SecretnessAnalysis.cpp b/lib/Analysis/SecretnessAnalysis/SecretnessAnalysis.cpp index 42284f1568..4f761c6c28 100644 --- a/lib/Analysis/SecretnessAnalysis/SecretnessAnalysis.cpp +++ b/lib/Analysis/SecretnessAnalysis/SecretnessAnalysis.cpp @@ -234,7 +234,7 @@ void annotateSecretness(Operation* top, DataFlowSolver* solver, bool verbose) { }); } -bool isSecret(Value value, DataFlowSolver* solver) { +bool isSecret(Value value, const DataFlowSolver* solver) { auto* lattice = solver->lookupState(value); return isSecret(lattice); } @@ -249,7 +249,7 @@ bool isSecret(const SecretnessLattice* lattice) { return lattice->getValue().getSecretness(); } -bool isSecret(ValueRange values, DataFlowSolver* solver) { +bool isSecret(ValueRange values, const DataFlowSolver* solver) { if (values.empty()) { return false; } diff --git a/lib/Analysis/SecretnessAnalysis/SecretnessAnalysis.h b/lib/Analysis/SecretnessAnalysis/SecretnessAnalysis.h index 28cc77f8be..7a511ea8b3 100644 --- a/lib/Analysis/SecretnessAnalysis/SecretnessAnalysis.h +++ b/lib/Analysis/SecretnessAnalysis/SecretnessAnalysis.h @@ -244,11 +244,11 @@ void annotateSecretness(Operation* top, DataFlowSolver* solver, bool verbose); // this method is used when DataFlowSolver has finished running the secretness // analysis -bool isSecret(Value value, DataFlowSolver* solver); +bool isSecret(Value value, const DataFlowSolver* solver); bool isSecret(const SecretnessLattice* lattice); -bool isSecret(ValueRange values, DataFlowSolver* solver); +bool isSecret(ValueRange values, const DataFlowSolver* solver); void getSecretOperands(Operation* op, SmallVectorImpl& secretOperands, diff --git a/lib/Dialect/Polynomial/IR/PolynomialOps.cpp b/lib/Dialect/Polynomial/IR/PolynomialOps.cpp index 6a7061d542..976b55fa30 100644 --- a/lib/Dialect/Polynomial/IR/PolynomialOps.cpp +++ b/lib/Dialect/Polynomial/IR/PolynomialOps.cpp @@ -1079,7 +1079,9 @@ OpFoldResult MulOp::fold(FoldAdaptor adaptor) { return nullptr; } - RNSPolynomial resultPoly = lhsPoly.mul(rhsPoly); + std::optional resultPolyOpt = lhsPoly.mul(rhsPoly); + if (!resultPolyOpt) return nullptr; + RNSPolynomial resultPoly = *resultPolyOpt; auto resultType = getResult().getType(); auto elementType = lhsAttr.getCoefficients().getElementType(); @@ -1102,7 +1104,9 @@ OpFoldResult MulOp::fold(FoldAdaptor adaptor) { auto lhs = getSingleLimbRNSPolynomial(lhsIntAttr, lhsPoly); auto rhs = getSingleLimbRNSPolynomial(rhsIntAttr, rhsPoly); if (lhs && rhs) { - RNSPolynomial result = lhs->mul(*rhs); + std::optional resultOpt = lhs->mul(*rhs); + if (!resultOpt) return nullptr; + RNSPolynomial result = *resultOpt; return getTypedIntPolynomialAttr(getContext(), result.getData(), getResult().getType()); } @@ -1126,7 +1130,9 @@ OpFoldResult NTTOp::fold(FoldAdaptor adaptor) { if (!rnsRootAttr) return nullptr; RNSPolynomial poly = inputAttr.getPolynomial(); - RNSPolynomial resultPoly = poly.toNtt(rnsRootAttr); + std::optional resultPolyOpt = poly.toNtt(rnsRootAttr); + if (!resultPolyOpt) return nullptr; + RNSPolynomial resultPoly = *resultPolyOpt; auto resultType = getResult().getType(); auto elementType = inputAttr.getCoefficients().getElementType(); @@ -1151,7 +1157,9 @@ OpFoldResult NTTOp::fold(FoldAdaptor adaptor) { if (!poly) return nullptr; SmallVector roots = { modArithRootAttr.getValue().getValue().getZExtValue()}; - RNSPolynomial resultPoly = poly->toNtt(roots); + std::optional resultPolyOpt = poly->toNtt(roots); + if (!resultPolyOpt) return nullptr; + RNSPolynomial resultPoly = *resultPolyOpt; return getTypedIntPolynomialAttr(getContext(), resultPoly.getData(), getResult().getType()); } @@ -1166,7 +1174,10 @@ OpFoldResult INTTOp::fold(FoldAdaptor adaptor) { if (!rnsRootAttr) return nullptr; RNSPolynomial poly = inputAttr.getPolynomial(); - RNSPolynomial resultPoly = poly.toCoefficient(rnsRootAttr); + std::optional resultPolyOpt = + poly.toCoefficient(rnsRootAttr); + if (!resultPolyOpt) return nullptr; + RNSPolynomial resultPoly = *resultPolyOpt; auto resultType = getResult().getType(); auto elementType = inputAttr.getCoefficients().getElementType(); @@ -1191,7 +1202,9 @@ OpFoldResult INTTOp::fold(FoldAdaptor adaptor) { if (!poly) return nullptr; SmallVector roots = { modArithRootAttr.getValue().getValue().getZExtValue()}; - RNSPolynomial resultPoly = poly->toCoefficient(roots); + std::optional resultPolyOpt = poly->toCoefficient(roots); + if (!resultPolyOpt) return nullptr; + RNSPolynomial resultPoly = *resultPolyOpt; return getTypedIntPolynomialAttr(getContext(), resultPoly.getData(), getResult().getType()); } diff --git a/lib/Target/Lattigo/BUILD b/lib/Target/Lattigo/BUILD index 5421ab7f22..b2cc5d53e6 100644 --- a/lib/Target/Lattigo/BUILD +++ b/lib/Target/Lattigo/BUILD @@ -31,6 +31,7 @@ cc_library( "@llvm-project//mlir:DialectUtils", "@llvm-project//mlir:FuncDialect", "@llvm-project//mlir:IR", + "@llvm-project//mlir:MathDialect", "@llvm-project//mlir:MemRefDialect", "@llvm-project//mlir:SCFDialect", "@llvm-project//mlir:Support", diff --git a/lib/Target/Lattigo/LattigoEmitter.cpp b/lib/Target/Lattigo/LattigoEmitter.cpp index e6d6893c7c..f5dfb3d8c5 100644 --- a/lib/Target/Lattigo/LattigoEmitter.cpp +++ b/lib/Target/Lattigo/LattigoEmitter.cpp @@ -148,6 +148,7 @@ LogicalResult LattigoEmitter::translate(Operation& op) { [&](auto op) { return printOperation(op); }) .Case( [&](auto op) { return printOperation(op); }) + .Case([&](auto op) { return printOperation(op); }) // Lattigo ops .Case< @@ -1059,6 +1060,28 @@ LogicalResult LattigoEmitter::printOperation(arith::XOrIOp op) { return printBinaryOp(op, op.getLhs(), op.getRhs(), "^"); } +LogicalResult LattigoEmitter::printOperation(math::SqrtOp op) { + imports.insert("\"math\""); + Type type = op.getOperand().getType(); + auto typeStringResult = convertType(type); + if (failed(typeStringResult)) return failure(); + std::string typeString = typeStringResult.value(); + + std::string operandName = getName(op.getOperand()); + std::string resultName = getName(op.getResult()); + + if (typeString == "float32") { + os << resultName << " := float32(math.Sqrt(float64(" << operandName + << ")))\n"; + } else if (typeString == "float64") { + os << resultName << " := math.Sqrt(" << operandName << ")\n"; + } else { + return op.emitOpError("Unsupported float type for math.sqrt: ") + << typeString; + } + return success(); +} + LogicalResult LattigoEmitter::printOperation(arith::RemSIOp op) { return printBinaryOp(op, op.getLhs(), op.getRhs(), "%"); } @@ -2604,7 +2627,7 @@ void registerToLattigoTranslation() { func::FuncDialect, tensor::TensorDialect, tensor_ext::TensorExtDialect, lattigo::LattigoDialect, memref::MemRefDialect, mgmt::MgmtDialect, scf::SCFDialect, - preprocessing::PreprocessingDialect>(); + preprocessing::PreprocessingDialect, math::MathDialect>(); }); } @@ -2625,7 +2648,7 @@ void registerToLattigoPreprocessingTranslation() { func::FuncDialect, tensor::TensorDialect, tensor_ext::TensorExtDialect, lattigo::LattigoDialect, memref::MemRefDialect, mgmt::MgmtDialect, scf::SCFDialect, - preprocessing::PreprocessingDialect>(); + preprocessing::PreprocessingDialect, math::MathDialect>(); }); } @@ -2646,7 +2669,7 @@ void registerToLattigoPreprocessedTranslation() { func::FuncDialect, tensor::TensorDialect, tensor_ext::TensorExtDialect, lattigo::LattigoDialect, memref::MemRefDialect, mgmt::MgmtDialect, scf::SCFDialect, - preprocessing::PreprocessingDialect>(); + preprocessing::PreprocessingDialect, math::MathDialect>(); }); } diff --git a/lib/Target/Lattigo/LattigoEmitter.h b/lib/Target/Lattigo/LattigoEmitter.h index 966135d3ec..086d44cace 100644 --- a/lib/Target/Lattigo/LattigoEmitter.h +++ b/lib/Target/Lattigo/LattigoEmitter.h @@ -18,6 +18,7 @@ #include "mlir/include/mlir/Dialect/Affine/IR/AffineOps.h" // from @llvm-project #include "mlir/include/mlir/Dialect/Arith/IR/Arith.h" // from @llvm-project #include "mlir/include/mlir/Dialect/Func/IR/FuncOps.h" // from @llvm-project +#include "mlir/include/mlir/Dialect/Math/IR/Math.h" // from @llvm-project #include "mlir/include/mlir/Dialect/MemRef/IR/MemRef.h" // from @llvm-project #include "mlir/include/mlir/Dialect/SCF/IR/SCF.h" // from @llvm-project #include "mlir/include/mlir/Dialect/Tensor/IR/Tensor.h" // from @llvm-project @@ -136,6 +137,7 @@ class LattigoEmitter { LogicalResult printOperation(::mlir::arith::SubIOp op); LogicalResult printOperation(::mlir::arith::SubFOp op); LogicalResult printOperation(::mlir::arith::XOrIOp op); + LogicalResult printOperation(::mlir::math::SqrtOp op); LogicalResult printOperation(::mlir::func::CallOp op); LogicalResult printOperation(::mlir::func::FuncOp op); LogicalResult printOperation(::mlir::func::ReturnOp op); diff --git a/lib/Transforms/LowerPolynomialEval/BUILD b/lib/Transforms/LowerPolynomialEval/BUILD index d4257e2985..ab23291eff 100644 --- a/lib/Transforms/LowerPolynomialEval/BUILD +++ b/lib/Transforms/LowerPolynomialEval/BUILD @@ -13,7 +13,11 @@ cc_library( deps = [ ":Patterns", ":pass_inc_gen", + "@heir//lib/Analysis/SecretnessAnalysis", + "@heir//lib/Dialect/Kernel/IR:Dialect", "@heir//lib/Dialect/Polynomial/IR:Dialect", + "@heir//lib/Target/CompilationTarget", + "@llvm-project//mlir:Analysis", "@llvm-project//mlir:IR", "@llvm-project//mlir:Pass", "@llvm-project//mlir:TransformUtils", @@ -26,6 +30,8 @@ cc_library( srcs = ["Patterns.cpp"], hdrs = ["Patterns.h"], deps = [ + "@heir//lib/Analysis/SecretnessAnalysis", + "@heir//lib/Dialect/Kernel/IR:Dialect", "@heir//lib/Dialect/Polynomial/IR:Dialect", "@heir//lib/Kernel:AbstractValue", "@heir//lib/Kernel:ArithmeticDag", diff --git a/lib/Transforms/LowerPolynomialEval/LowerPolynomialEval.cpp b/lib/Transforms/LowerPolynomialEval/LowerPolynomialEval.cpp index 29b74f8618..4611a754e8 100644 --- a/lib/Transforms/LowerPolynomialEval/LowerPolynomialEval.cpp +++ b/lib/Transforms/LowerPolynomialEval/LowerPolynomialEval.cpp @@ -2,9 +2,13 @@ #include +#include "lib/Analysis/SecretnessAnalysis/SecretnessAnalysis.h" +#include "lib/Dialect/Kernel/IR/KernelDialect.h" +#include "lib/Target/CompilationTarget/CompilationTarget.h" #include "lib/Transforms/LowerPolynomialEval/Patterns.h" -#include "mlir/include/mlir/IR/MLIRContext.h" // from @llvm-project -#include "mlir/include/mlir/IR/PatternMatch.h" // from @llvm-project +#include "mlir/include/mlir/Analysis/DataFlow/Utils.h" // from @llvm-project +#include "mlir/include/mlir/IR/MLIRContext.h" // from @llvm-project +#include "mlir/include/mlir/IR/PatternMatch.h" // from @llvm-project #include "mlir/include/mlir/Transforms/WalkPatternRewriteDriver.h" // from @llvm-project // IWYU pragma: begin_keep @@ -17,21 +21,62 @@ namespace heir { #define GEN_PASS_DEF_LOWERPOLYNOMIALEVAL #include "lib/Transforms/LowerPolynomialEval/LowerPolynomialEval.h.inc" +static bool hasBackendAttribute(ModuleOp module) { + if (!module) return false; + for (NamedAttribute attr : module->getAttrs()) { + if (!isa(attr.getValue())) continue; + if (attr.getName().strref().starts_with("backend.")) { + return true; + } + } + return false; +} + struct LowerPolynomialEval : impl::LowerPolynomialEvalBase { using LowerPolynomialEvalBase::LowerPolynomialEvalBase; void runOnOperation() override { MLIRContext* context = &getContext(); + + ModuleOp module = dyn_cast(getOperation()); + if (!module) { + module = getOperation()->getParentOfType(); + } + + bool hasKernelChebyshev = false; + if (module && hasBackendAttribute(module)) { + auto target = getTargetConfig(module); + if (succeeded(target)) { + hasKernelChebyshev = target->has_kernel_chebyshev; + } + } + RewritePatternSet patterns(context); + DataFlowSolver solver; + dataflow::loadBaselineAnalyses(solver); + solver.load(); + if (failed(solver.initializeAndRun(getOperation()))) { + getOperation()->emitOpError() << "Failed to run SecretnessAnalysis.\n"; + return signalPassFailure(); + } + switch (method) { case PolynomialApproximationMethod::Automatic: patterns.add( context, /*force=*/false); - patterns.add( - context, - /*force=*/false, minCoefficientThreshold); + if (hasKernelChebyshev) { + patterns.add(context, solver, + /*force=*/false); + patterns.add( + context, + /*force=*/false, minCoefficientThreshold); + } else { + patterns.add( + context, + /*force=*/false, minCoefficientThreshold); + } break; case PolynomialApproximationMethod::Horner: patterns.add(context, /*force=*/true); @@ -41,9 +86,17 @@ struct LowerPolynomialEval /*force=*/true); break; case PolynomialApproximationMethod::PatersonStockmeyerChebyshev: - patterns.add( - context, - /*force=*/true, minCoefficientThreshold); + if (hasKernelChebyshev) { + patterns.add(context, solver, + /*force=*/true); + patterns.add( + context, + /*force=*/true, minCoefficientThreshold); + } else { + patterns.add( + context, + /*force=*/true, minCoefficientThreshold); + } break; default: getOperation()->emitError() << "Unknown lowering method: " << method; diff --git a/lib/Transforms/LowerPolynomialEval/LowerPolynomialEval.td b/lib/Transforms/LowerPolynomialEval/LowerPolynomialEval.td index 7d04ca7460..3d0ebf181e 100644 --- a/lib/Transforms/LowerPolynomialEval/LowerPolynomialEval.td +++ b/lib/Transforms/LowerPolynomialEval/LowerPolynomialEval.td @@ -29,6 +29,7 @@ def LowerPolynomialEval : Pass<"lower-polynomial-eval"> { }]; let dependentDialects = [ "::mlir::heir::polynomial::PolynomialDialect", + "::mlir::heir::kernel::KernelDialect", ]; let options = [ Option<"method", "method", "mlir::heir::PolynomialApproximationMethod", diff --git a/lib/Transforms/LowerPolynomialEval/Patterns.cpp b/lib/Transforms/LowerPolynomialEval/Patterns.cpp index 5b9bed9a3b..adce8d9d57 100644 --- a/lib/Transforms/LowerPolynomialEval/Patterns.cpp +++ b/lib/Transforms/LowerPolynomialEval/Patterns.cpp @@ -4,6 +4,8 @@ #include #include +#include "lib/Analysis/SecretnessAnalysis/SecretnessAnalysis.h" +#include "lib/Dialect/Kernel/IR/KernelOps.h" #include "lib/Dialect/Polynomial/IR/PolynomialAttributes.h" #include "lib/Dialect/Polynomial/IR/PolynomialOps.h" #include "lib/Kernel/AbstractValue.h" @@ -226,5 +228,51 @@ LogicalResult LowerViaPatersonStockmeyerChebyshev::matchAndRewrite( return success(); } +LogicalResult LowerToKernelEvalChebyshev::matchAndRewrite( + EvalOp op, PatternRewriter& rewriter) const { + if (!mlir::heir::isSecret(op.getValue(), &solver)) { + return rewriter.notifyMatchFailure(op, "operand is not secret"); + } + auto attr = dyn_cast( + op.getPolynomialAttr()); + if (!attr) return failure(); + + auto lowerAttr = op->getAttrOfType("domain_lower"); + auto upperAttr = op->getAttrOfType("domain_upper"); + if (!lowerAttr || !upperAttr) return failure(); + + double lower = lowerAttr.getValue().convertToDouble(); + double upper = upperAttr.getValue().convertToDouble(); + + ImplicitLocOpBuilder b(op.getLoc(), rewriter); + Value xInput = op.getValue(); + + if (std::abs(lower - -1.0) > 1e-9 || std::abs(upper - 1.0) > 1e-9) { + APFloat rescale = APFloat(2.0 / (upper - lower)); + APFloat shift = APFloat(-(upper + lower) / (upper - lower)); + + Type inputType = xInput.getType(); + + if (!rescale.isExactlyValue(1.0)) { + xInput = arith::MulFOp::create( + b, xInput, + arith::ConstantOp::create( + b, inputType, getScalarOrDenseAttr(inputType, rescale))) + .getResult(); + } + if (!shift.isZero()) { + xInput = arith::AddFOp::create( + b, xInput, + arith::ConstantOp::create( + b, inputType, getScalarOrDenseAttr(inputType, shift))) + .getResult(); + } + } + + rewriter.replaceOpWithNewOp( + op, op.getType(), xInput, attr.getValue().getCoefficients()); + return success(); +} + } // namespace heir } // namespace mlir diff --git a/lib/Transforms/LowerPolynomialEval/Patterns.h b/lib/Transforms/LowerPolynomialEval/Patterns.h index efc8d59626..228c685867 100644 --- a/lib/Transforms/LowerPolynomialEval/Patterns.h +++ b/lib/Transforms/LowerPolynomialEval/Patterns.h @@ -9,11 +9,14 @@ // Lowering patterns for polynomial.eval. namespace mlir { +class DataFlowSolver; namespace heir { struct LoweringBase : public OpRewritePattern { - LoweringBase(MLIRContext* context, bool force = false) - : mlir::OpRewritePattern(context), force(force) {} + LoweringBase(MLIRContext* context, bool force = false, + PatternBenefit benefit = 1) + : mlir::OpRewritePattern(context, benefit), + force(force) {} bool shouldForce() const { return force; } @@ -68,6 +71,20 @@ struct LowerViaPatersonStockmeyerChebyshev : public ChebyshevLoweringBase { PatternRewriter& rewriter) const override; }; +// Lower polynomial.eval that uses a Chebyshev float polynomial to +// kernel.eval_chebyshev. +struct LowerToKernelEvalChebyshev : public LoweringBase { + LowerToKernelEvalChebyshev(MLIRContext* context, const DataFlowSolver& solver, + bool force = false) + : LoweringBase(context, force, /*benefit=*/2), solver(solver) {} + + LogicalResult matchAndRewrite(polynomial::EvalOp op, + PatternRewriter& rewriter) const override; + + private: + const DataFlowSolver& solver; +}; + } // namespace heir } // namespace mlir diff --git a/lib/Transforms/PolynomialApproximation/BUILD b/lib/Transforms/PolynomialApproximation/BUILD index 63add27bd4..5c2e00820d 100644 --- a/lib/Transforms/PolynomialApproximation/BUILD +++ b/lib/Transforms/PolynomialApproximation/BUILD @@ -12,11 +12,13 @@ cc_library( hdrs = ["PolynomialApproximation.h"], deps = [ ":pass_inc_gen", + "@heir//lib/Analysis/SecretnessAnalysis", "@heir//lib/Dialect/MathExt/IR:Dialect", "@heir//lib/Dialect/Polynomial/IR:Dialect", "@heir//lib/Utils/Approximation:CaratheodoryFejer", "@heir//lib/Utils/Polynomial", "@llvm-project//llvm:Support", + "@llvm-project//mlir:Analysis", "@llvm-project//mlir:ArithDialect", "@llvm-project//mlir:IR", "@llvm-project//mlir:MathDialect", diff --git a/lib/Transforms/PolynomialApproximation/PolynomialApproximation.cpp b/lib/Transforms/PolynomialApproximation/PolynomialApproximation.cpp index 574129e049..8821e93ca2 100644 --- a/lib/Transforms/PolynomialApproximation/PolynomialApproximation.cpp +++ b/lib/Transforms/PolynomialApproximation/PolynomialApproximation.cpp @@ -5,17 +5,19 @@ #include #include +#include "lib/Analysis/SecretnessAnalysis/SecretnessAnalysis.h" #include "lib/Dialect/MathExt/IR/MathExtOps.h" #include "lib/Dialect/Polynomial/IR/PolynomialAttributes.h" #include "lib/Dialect/Polynomial/IR/PolynomialOps.h" #include "lib/Dialect/Polynomial/IR/PolynomialTypes.h" #include "lib/Utils/Approximation/CaratheodoryFejer.h" #include "lib/Utils/Polynomial/Polynomial.h" -#include "llvm/include/llvm/ADT/APFloat.h" // from @llvm-project -#include "llvm/include/llvm/Support/Casting.h" // from @llvm-project -#include "llvm/include/llvm/Support/Debug.h" // from @llvm-project -#include "mlir/include/mlir/Dialect/Arith/IR/Arith.h" // from @llvm-project -#include "mlir/include/mlir/Dialect/Math/IR/Math.h" // from @llvm-project +#include "llvm/include/llvm/ADT/APFloat.h" // from @llvm-project +#include "llvm/include/llvm/Support/Casting.h" // from @llvm-project +#include "llvm/include/llvm/Support/Debug.h" // from @llvm-project +#include "mlir/include/mlir/Analysis/DataFlow/Utils.h" // from @llvm-project +#include "mlir/include/mlir/Dialect/Arith/IR/Arith.h" // from @llvm-project +#include "mlir/include/mlir/Dialect/Math/IR/Math.h" // from @llvm-project #include "mlir/include/mlir/IR/BuiltinAttributeInterfaces.h" // from @llvm-project #include "mlir/include/mlir/IR/BuiltinAttributes.h" // from @llvm-project #include "mlir/include/mlir/IR/BuiltinTypeInterfaces.h" // from @llvm-project @@ -214,11 +216,12 @@ inline APFloat minnumf(const APFloat& lhs, const APFloat& rhs) { template struct ConvertUnaryOp : public OpRewritePattern { - ConvertUnaryOp(mlir::MLIRContext* context, + ConvertUnaryOp(mlir::MLIRContext* context, DataFlowSolver* solver, const std::function& cppFunc, double lower = kDefaultDomainLower, double upper = kDefaultDomainUpper) : OpRewritePattern(context, /*benefit=*/1), + solver(solver), cppFunc(cppFunc), lower(lower), upper(upper) {} @@ -226,6 +229,9 @@ struct ConvertUnaryOp : public OpRewritePattern { public: LogicalResult matchAndRewrite(OpTy op, PatternRewriter& rewriter) const override { + if (!mlir::heir::isSecret(op.getOperand(), solver)) { + return rewriter.notifyMatchFailure(op, "operand is not secret"); + } MLIRContext* ctx = op.getContext(); IntegerAttr degreeAttr = op->hasAttr("degree") ? cast(op->getAttr("degree")) @@ -266,6 +272,7 @@ struct ConvertUnaryOp : public OpRewritePattern { } private: + DataFlowSolver* solver; std::function cppFunc; double lower; double upper; @@ -303,11 +310,12 @@ FailureOr getSingleValueOrSplat(Value value) { template struct ConvertBinaryConstOp : public OpRewritePattern { - ConvertBinaryConstOp(mlir::MLIRContext* context, + ConvertBinaryConstOp(mlir::MLIRContext* context, DataFlowSolver* solver, const std::function& cppFunc, double lower = kDefaultDomainLower, double upper = kDefaultDomainUpper) : OpRewritePattern(context, /*benefit=*/1), + solver(solver), cppFunc(cppFunc), lower(lower), upper(upper) {} @@ -337,6 +345,10 @@ struct ConvertBinaryConstOp : public OpRewritePattern { lhsIsConstant ? lhsConstResult.value() : rhsConstResult.value(); Value nonConstOperand = lhsIsConstant ? rhs : lhs; + if (!mlir::heir::isSecret(nonConstOperand, solver)) { + return rewriter.notifyMatchFailure(op, "operand is not secret"); + } + // cppFunc is a binary op, so we need to give it the constant value to // convert it to a unary op. std::function unaryFunc; @@ -389,6 +401,7 @@ struct ConvertBinaryConstOp : public OpRewritePattern { } private: + DataFlowSolver* solver; std::function cppFunc; double lower; double upper; @@ -398,14 +411,19 @@ struct ConvertBinaryConstOp : public OpRewritePattern { // repeated squaring. When the domain is in [-2^k, 1], this is more efficient // in level consumption than the default polynomial approximation solver. struct ExpOpTaylorApproximation : public OpRewritePattern { - ExpOpTaylorApproximation(MLIRContext* context, int64_t defaultK = 7) + ExpOpTaylorApproximation(MLIRContext* context, DataFlowSolver* solver, + int64_t defaultK = 7) : OpRewritePattern(context, /*benefit=*/2), + solver(solver), defaultK(defaultK) {} LogicalResult matchAndRewrite(math::ExpOp op, PatternRewriter& rewriter) const override { Location loc = op.getLoc(); Value operand = op.getOperand(); + if (!mlir::heir::isSecret(operand, solver)) { + return rewriter.notifyMatchFailure(op, "operand is not secret"); + } Type type = operand.getType(); int64_t k = defaultK; @@ -473,6 +491,7 @@ struct ExpOpTaylorApproximation : public OpRewritePattern { } private: + DataFlowSolver* solver; int64_t defaultK; }; @@ -482,52 +501,67 @@ struct PolynomialApproximation void runOnOperation() override { MLIRContext* context = &getContext(); + + DataFlowSolver solver; + dataflow::loadBaselineAnalyses(solver); + solver.load(); + if (failed(solver.initializeAndRun(getOperation()))) { + getOperation()->emitOpError() << "Failed to run SecretnessAnalysis.\n"; + return signalPassFailure(); + } + RewritePatternSet patterns(context); // High priority patterns - patterns.add(context, /*k=*/7); + patterns.add(context, &solver, /*k=*/7); // Math unary ops - patterns.add>(context, absf); - patterns.add>(context, acos); - patterns.add>(context, acosh); - patterns.add>(context, asin); - patterns.add>(context, asinh); - patterns.add>(context, atan); - patterns.add>(context, atanh); - patterns.add>(context, cbrt); - patterns.add>(context, ceil); - patterns.add>(context, cos); - patterns.add>(context, cosh); - patterns.add>(context, erf); - patterns.add>(context, erfc); - patterns.add>(context, exp); - patterns.add>(context, exp2); - patterns.add>(context, expm1); - patterns.add>(context, floor); - patterns.add>( - context, log, kDefaultPositiveRangeLower, kDefaultPositiveRangeUpper); - patterns.add>( - context, log10, kDefaultPositiveRangeLower, kDefaultPositiveRangeUpper); - patterns.add>(context, log1p); - patterns.add>( - context, log2, kDefaultPositiveRangeLower, kDefaultPositiveRangeUpper); - patterns.add>(context, round); - patterns.add>( - context, rsqrt, kDefaultPositiveRangeLower, kDefaultPositiveRangeUpper); - patterns.add>(context, sin); - patterns.add>(context, sinh); - patterns.add>(context, sqrt, + patterns.add>(context, &solver, absf); + patterns.add>(context, &solver, acos); + patterns.add>(context, &solver, acosh); + patterns.add>(context, &solver, asin); + patterns.add>(context, &solver, asinh); + patterns.add>(context, &solver, atan); + patterns.add>(context, &solver, atanh); + patterns.add>(context, &solver, cbrt); + patterns.add>(context, &solver, ceil); + patterns.add>(context, &solver, cos); + patterns.add>(context, &solver, cosh); + patterns.add>(context, &solver, erf); + patterns.add>(context, &solver, erfc); + patterns.add>(context, &solver, exp); + patterns.add>(context, &solver, exp2); + patterns.add>(context, &solver, expm1); + patterns.add>(context, &solver, floor); + patterns.add>(context, &solver, log, + kDefaultPositiveRangeLower, + kDefaultPositiveRangeUpper); + patterns.add>(context, &solver, log10, + kDefaultPositiveRangeLower, + kDefaultPositiveRangeUpper); + patterns.add>(context, &solver, log1p); + patterns.add>(context, &solver, log2, + kDefaultPositiveRangeLower, + kDefaultPositiveRangeUpper); + patterns.add>(context, &solver, round); + patterns.add>(context, &solver, rsqrt, + kDefaultPositiveRangeLower, + kDefaultPositiveRangeUpper); + patterns.add>(context, &solver, sin); + patterns.add>(context, &solver, sinh); + patterns.add>(context, &solver, sqrt, kDefaultNonNegativeRangeLower, kDefaultNonNegativeRangeUpper); - patterns.add>(context, tan); - patterns.add>(context, tanh); - patterns.add>(context, trunc); - patterns.add>(context, sign); - patterns.add>(context, sigmoid); + patterns.add>(context, &solver, tan); + patterns.add>(context, &solver, tanh); + patterns.add>(context, &solver, trunc); + patterns.add>(context, &solver, sign); + patterns.add>(context, &solver, + sigmoid); // TODO(#1514): Restore with alternative roundeven - // patterns.add>(context, _roundeven); + // patterns.add>(context, &solver, + // _roundeven); // Unsupported math dialect unary ops: // math::AbsIOp @@ -540,17 +574,22 @@ struct PolynomialApproximation // math::IsnormalOp // Math binary ops (when one argument is statically constant) - patterns.add>(context, maxnumf); - patterns.add>(context, maxf); - patterns.add>(context, minf); - patterns.add>(context, minnumf); - patterns.add>(context, atan2); - patterns.add>(context, copysign); - patterns.add>(context, fpowi); - patterns.add>(context, powf); + patterns.add>(context, &solver, + maxnumf); + patterns.add>(context, &solver, + maxf); + patterns.add>(context, &solver, + minf); + patterns.add>(context, &solver, + minnumf); + patterns.add>(context, &solver, atan2); + patterns.add>(context, &solver, + copysign); + patterns.add>(context, &solver, fpowi); + patterns.add>(context, &solver, powf); // Math ternary ops - // patterns.add>(context, fma); + // patterns.add>(context, &solver, fma); // TODO (#1221): Investigate whether folding (default: on) can be skipped // here. diff --git a/lib/Utils/Polynomial/PolynomialTest.cpp b/lib/Utils/Polynomial/PolynomialTest.cpp index 22fdea453c..229663ebc9 100644 --- a/lib/Utils/Polynomial/PolynomialTest.cpp +++ b/lib/Utils/Polynomial/PolynomialTest.cpp @@ -216,10 +216,10 @@ TEST(RNSPolynomialTest, TestConversions) { RNSPolynomial poly(coeffs, moduli); // Test round-trip toNtt -> toCoefficient - RNSPolynomial ntt = poly.toNtt(); + RNSPolynomial ntt = poly.toNtt().value(); EXPECT_TRUE(ntt.isNtt()); - RNSPolynomial roundtrip = ntt.toCoefficient(); + RNSPolynomial roundtrip = ntt.toCoefficient().value(); EXPECT_FALSE(roundtrip.isNtt()); EXPECT_EQ(roundtrip, poly); } @@ -233,17 +233,17 @@ TEST(RNSPolynomialTest, TestMul) { RNSPolynomial poly2(coeffs2, moduli); // Test multiplication in Coefficient form (uses NTT under the hood) - RNSPolynomial prodCoeff = poly1.mul(poly2); + RNSPolynomial prodCoeff = poly1.mul(poly2).value(); EXPECT_FALSE(prodCoeff.isNtt()); SmallVector expectedProd = {5, 16, 12, 0, 21, 11, 32, 0}; EXPECT_EQ(prodCoeff.getData(), llvm::ArrayRef(expectedProd)); // Test multiplication in NTT form - RNSPolynomial ntt1 = poly1.toNtt(); - RNSPolynomial ntt2 = poly2.toNtt(); - RNSPolynomial prodNtt = ntt1.mul(ntt2); + RNSPolynomial ntt1 = poly1.toNtt().value(); + RNSPolynomial ntt2 = poly2.toNtt().value(); + RNSPolynomial prodNtt = ntt1.mul(ntt2).value(); EXPECT_TRUE(prodNtt.isNtt()); - EXPECT_EQ(prodNtt.toCoefficient(), prodCoeff); + EXPECT_EQ(prodNtt.toCoefficient().value(), prodCoeff); } TEST(RNSPolynomialTest, TestPrecomputedRoots) { @@ -283,15 +283,15 @@ TEST(RNSPolynomialTest, TestPrecomputedRoots) { mlir::heir::rns::RNSAttr::get(rnsType, {root17Attr, root41Attr}); // Test toNtt with precomputed roots (matching on-the-fly) - RNSPolynomial ntt = poly.toNtt(rnsAttr); + RNSPolynomial ntt = poly.toNtt(rnsAttr).value(); EXPECT_TRUE(ntt.isNtt()); // Compare with on-the-fly computation - RNSPolynomial nttOnTheFly = poly.toNtt(); + RNSPolynomial nttOnTheFly = poly.toNtt().value(); EXPECT_EQ(ntt, nttOnTheFly); // Test toCoefficient with precomputed roots - RNSPolynomial roundtrip = ntt.toCoefficient(rnsAttr); + RNSPolynomial roundtrip = ntt.toCoefficient(rnsAttr).value(); EXPECT_FALSE(roundtrip.isNtt()); EXPECT_EQ(roundtrip, poly); @@ -305,10 +305,10 @@ TEST(RNSPolynomialTest, TestPrecomputedRoots) { auto diffRnsAttr = mlir::heir::rns::RNSAttr::get(rnsType, {diffRoot17Attr, diffRoot41Attr}); - RNSPolynomial nttDiff = poly.toNtt(diffRnsAttr); + RNSPolynomial nttDiff = poly.toNtt(diffRnsAttr).value(); EXPECT_TRUE(nttDiff.isNtt()); - RNSPolynomial roundtripDiff = nttDiff.toCoefficient(diffRnsAttr); + RNSPolynomial roundtripDiff = nttDiff.toCoefficient(diffRnsAttr).value(); EXPECT_FALSE(roundtripDiff.isNtt()); EXPECT_EQ(roundtripDiff, poly); } diff --git a/lib/Utils/Polynomial/RNSPolynomial.cpp b/lib/Utils/Polynomial/RNSPolynomial.cpp index d0cc5baa3a..7a589c353d 100644 --- a/lib/Utils/Polynomial/RNSPolynomial.cpp +++ b/lib/Utils/Polynomial/RNSPolynomial.cpp @@ -97,7 +97,8 @@ std::optional RNSPolynomial::scalarMul( return RNSPolynomial(std::move(resultData), moduli, representation); } -RNSPolynomial RNSPolynomial::mul(const RNSPolynomial& other) const { +std::optional RNSPolynomial::mul( + const RNSPolynomial& other) const { assert(moduli == other.moduli && "Moduli must match for multiplication"); assert(numCoeffs == other.numCoeffs && "Number of coefficients must match"); assert(representation == other.representation && @@ -121,14 +122,23 @@ RNSPolynomial RNSPolynomial::mul(const RNSPolynomial& other) const { } if (representation == Form::COEFF && other.representation == Form::COEFF) { - return toNtt().mul(other.toNtt()).toCoefficient(); + auto lhsNtt = toNtt(); + if (!lhsNtt) return std::nullopt; + auto rhsNtt = other.toNtt(); + if (!rhsNtt) return std::nullopt; + auto prodNtt = lhsNtt->mul(*rhsNtt); + if (!prodNtt) return std::nullopt; + return prodNtt->toCoefficient(); } - return RNSPolynomial(); + return std::nullopt; } -RNSPolynomial RNSPolynomial::toNtt( +std::optional RNSPolynomial::toNtt( llvm::ArrayRef rootsOfUnity) const { + if (numCoeffs > 0 && (numCoeffs & (numCoeffs - 1)) != 0) { + return std::nullopt; + } assert(representation == Form::COEFF && "Already in NTT representation"); llvm::SmallVector resultData; @@ -149,7 +159,7 @@ RNSPolynomial RNSPolynomial::toNtt( llvm::APInt qAp(64, modulus); std::optional rootOpt = findPrimitive2nthRoot(qAp, numCoeffs); - assert(rootOpt.has_value() && "Primitive 2n-th root of unity not found"); + if (!rootOpt.has_value()) return std::nullopt; rootOfUnity = rootOpt->getZExtValue(); } @@ -166,7 +176,7 @@ RNSPolynomial RNSPolynomial::toNtt( return RNSPolynomial(std::move(resultData), moduli, Form::EVAL); } -RNSPolynomial RNSPolynomial::toNtt(rns::RNSAttr rootAttr) const { +std::optional RNSPolynomial::toNtt(rns::RNSAttr rootAttr) const { if (!rootAttr) return toNtt(llvm::ArrayRef()); llvm::SmallVector roots; @@ -178,8 +188,11 @@ RNSPolynomial RNSPolynomial::toNtt(rns::RNSAttr rootAttr) const { return toNtt(roots); } -RNSPolynomial RNSPolynomial::toCoefficient( +std::optional RNSPolynomial::toCoefficient( llvm::ArrayRef rootsOfUnity) const { + if (numCoeffs > 0 && (numCoeffs & (numCoeffs - 1)) != 0) { + return std::nullopt; + } assert(representation == Form::EVAL && "Already in Coefficient representation"); @@ -201,7 +214,7 @@ RNSPolynomial RNSPolynomial::toCoefficient( llvm::APInt qAp(64, modulus); std::optional rootOpt = findPrimitive2nthRoot(qAp, numCoeffs); - assert(rootOpt.has_value() && "Primitive 2n-th root of unity not found"); + if (!rootOpt.has_value()) return std::nullopt; rootOfUnity = rootOpt->getZExtValue(); } @@ -218,7 +231,8 @@ RNSPolynomial RNSPolynomial::toCoefficient( return RNSPolynomial(std::move(resultData), moduli, Form::COEFF); } -RNSPolynomial RNSPolynomial::toCoefficient(rns::RNSAttr rootAttr) const { +std::optional RNSPolynomial::toCoefficient( + rns::RNSAttr rootAttr) const { if (!rootAttr) return toCoefficient(llvm::ArrayRef()); llvm::SmallVector roots; diff --git a/lib/Utils/Polynomial/RNSPolynomial.h b/lib/Utils/Polynomial/RNSPolynomial.h index b21989b4f5..ffc1001a1a 100644 --- a/lib/Utils/Polynomial/RNSPolynomial.h +++ b/lib/Utils/Polynomial/RNSPolynomial.h @@ -74,15 +74,18 @@ class RNSPolynomial { /// Performs modular multiplication limb-wise. In NTT form, this corresponds /// to an elementwise product. In coefficient form, it first converts to NTT /// form, multiplies elementwise, and converts back to coefficient form. - RNSPolynomial mul(const RNSPolynomial& other) const; + std::optional mul(const RNSPolynomial& other) const; /// Convert the polynomial to NTT representation. - RNSPolynomial toNtt(llvm::ArrayRef rootOfUnity) const; - RNSPolynomial toNtt(rns::RNSAttr rootAttr = nullptr) const; + std::optional toNtt( + llvm::ArrayRef rootOfUnity) const; + std::optional toNtt(rns::RNSAttr rootAttr = nullptr) const; /// Convert the polynomial to Coefficient representation. - RNSPolynomial toCoefficient(llvm::ArrayRef rootOfUnity) const; - RNSPolynomial toCoefficient(rns::RNSAttr rootAttr = nullptr) const; + std::optional toCoefficient( + llvm::ArrayRef rootOfUnity) const; + std::optional toCoefficient( + rns::RNSAttr rootAttr = nullptr) const; /// Slice the polynomial's RNS basis. RNSPolynomial slice(size_t start, size_t size) const; diff --git a/tests/Examples/orion/chebyshev/chebyshev.mlir b/tests/Examples/orion/chebyshev/chebyshev.mlir index 40fb7ede16..86d4f09187 100644 --- a/tests/Examples/orion/chebyshev/chebyshev.mlir +++ b/tests/Examples/orion/chebyshev/chebyshev.mlir @@ -18,7 +18,7 @@ #ciphertext_space_L10 = #lwe.ciphertext_space !ct_L10 = !lwe.lwe_ciphertext, ciphertext_space = #ciphertext_space_L10, key = #key, modulus_chain = #modulus_chain_L10_C10> module attributes {scheme.ckks, ckks.schemeParam = #ckks.scheme_param} { - func.func @chebyshev(%ct: !ct_L10) -> !ct_L10 { + func.func @chebyshev(%ct: !ct_L10 {secret.secret}) -> !ct_L10 { %ct_0 = orion.chebyshev %ct {coefficients = [0.0, 0.75, 0.0, 0.25], domain_end = 1.000000e+00 : f64, domain_start = -1.000000e+00 : f64} : (!ct_L10) -> !ct_L10 return %ct_0 : !ct_L10 } diff --git a/tests/Pipelines/math_to_polynomial_approximation/polynomial_approximation.mlir b/tests/Pipelines/math_to_polynomial_approximation/polynomial_approximation.mlir index cd28ee1d72..8d26770c37 100644 --- a/tests/Pipelines/math_to_polynomial_approximation/polynomial_approximation.mlir +++ b/tests/Pipelines/math_to_polynomial_approximation/polynomial_approximation.mlir @@ -1,7 +1,7 @@ // RUN: heir-opt --math-to-polynomial-approximation %s | FileCheck %s --dump-input=always // CHECK: @test_maximumf -func.func @test_maximumf(%x: tensor<10xf32>) -> tensor<10xf32> { +func.func @test_maximumf(%x: tensor<10xf32> {secret.secret}) -> tensor<10xf32> { // CHECK-NOT: arith.maximumf // CHECK-NOT: polynomial.eval diff --git a/tests/Regression/issue_2888.mlir b/tests/Regression/issue_2888.mlir index 081b2fadde..acd8cfc155 100644 --- a/tests/Regression/issue_2888.mlir +++ b/tests/Regression/issue_2888.mlir @@ -9,7 +9,7 @@ // RUN: heir-opt %s --polynomial-approximation --lower-polynomial-eval --verify-diagnostics --split-input-file !poly_ty = !polynomial.polynomial> -func.func @monomial_nan_direct(%x: f64) -> f64 { +func.func @monomial_nan_direct(%x: f64 {secret.secret}) -> f64 { // expected-error@+1 {{non-finite}} %0 = polynomial.eval #polynomial.typed_float_polynomial< 0x7FF8000000000000 @@ -22,7 +22,7 @@ func.func @monomial_nan_direct(%x: f64) -> f64 { // ----- !poly_ty = !polynomial.polynomial> -func.func @chebyshev_nan_direct(%x: f64) -> f64 { +func.func @chebyshev_nan_direct(%x: f64 {secret.secret}) -> f64 { // expected-error@+1 {{non-finite}} %0 = polynomial.eval #polynomial.typed_chebyshev_polynomial<[ 0x7FF8000000000000 : f64, 0x7FF8000000000000 : f64, @@ -33,14 +33,14 @@ func.func @chebyshev_nan_direct(%x: f64) -> f64 { } // ----- -func.func @sqrt_on_negative_domain(%x: f32) -> f32 { +func.func @sqrt_on_negative_domain(%x: f32 {secret.secret}) -> f32 { // expected-error@+1 {{non-finite}} %0 = math.sqrt %x {domain_lower = -1.0 : f64, domain_upper = 1.0 : f64} : f32 return %0 : f32 } // ----- -func.func @sqrt_on_positive_domain(%x: f32) -> f32 { +func.func @sqrt_on_positive_domain(%x: f32 {secret.secret}) -> f32 { // CHECK-NOT: math.sqrt %0 = math.sqrt %x {domain_lower = 0.25 : f64, domain_upper = 4.0 : f64} : f32 return %0 : f32 diff --git a/tests/Transforms/lower_polynomial_eval/conditional_lower.mlir b/tests/Transforms/lower_polynomial_eval/conditional_lower.mlir new file mode 100644 index 0000000000..0240c96a64 --- /dev/null +++ b/tests/Transforms/lower_polynomial_eval/conditional_lower.mlir @@ -0,0 +1,60 @@ +// RUN: heir-opt %s --lower-polynomial-eval --split-input-file | FileCheck %s + +// ----- + +module attributes { + backend.lattigo, + backend.config_override = {has_kernel_chebyshev = true} +} { + // CHECK: @test_secret_kernel + func.func @test_secret_kernel(%x: !secret.secret) -> !secret.secret { + // CHECK: secret.generic + // CHECK: kernel.eval_chebyshev %{{.*}} {coefficients = [1.000000e+00, 2.000000e+00]} : f64 -> f64 + %0 = secret.generic(%x : !secret.secret) { + ^body(%x_val: f64): + %1 = polynomial.eval #polynomial.typed_chebyshev_polynomial<[1.0, 2.0]> : !polynomial.polynomial>, %x_val {domain_lower = -1.0 : f64, domain_upper = 1.0 : f64} : f64 + secret.yield %1 : f64 + } -> (!secret.secret) + return %0 : !secret.secret + } + + // CHECK: @test_public_kernel + func.func @test_public_kernel(%x: f64) -> f64 { + // CHECK-NOT: kernel.eval_chebyshev + // CHECK: arith.addf + %0 = polynomial.eval #polynomial.typed_chebyshev_polynomial<[1.0, 2.0]> : !polynomial.polynomial>, %x {domain_lower = -1.0 : f64, domain_upper = 1.0 : f64} : f64 + return %0 : f64 + } + + // CHECK: @eval_chebyshev_scaling + func.func @eval_chebyshev_scaling(%x: !secret.secret) -> !secret.secret { + // CHECK: secret.generic + // CHECK-NEXT: ^body(%[[VAL_0:.*]]: f64): + // CHECK: %[[CST_0:.*]] = arith.constant 5.000000e-01 : f64 + // CHECK: %[[VAL_1:.*]] = arith.mulf %[[VAL_0]], %[[CST_0]] : f64 + // CHECK: %[[CST_1:.*]] = arith.constant -1.000000e+00 : f64 + // CHECK: %[[VAL_2:.*]] = arith.addf %[[VAL_1]], %[[CST_1]] : f64 + // CHECK: kernel.eval_chebyshev %[[VAL_2]] {coefficients = [1.000000e+00, 2.000000e+00]} : f64 -> f64 + %0 = secret.generic(%x : !secret.secret) { + ^body(%x_val: f64): + %1 = polynomial.eval #polynomial.typed_chebyshev_polynomial<[1.0, 2.0]> : !polynomial.polynomial>, %x_val {domain_lower = 0.0 : f64, domain_upper = 4.0 : f64} : f64 + secret.yield %1 : f64 + } -> (!secret.secret) + return %0 : !secret.secret + } +} + +// ----- + +module attributes { + backend.lattigo, + backend.config_override = {has_kernel_chebyshev = false} +} { + // CHECK: @test_arith + func.func @test_arith(%x: f64) -> f64 { + // CHECK-NOT: kernel.eval_chebyshev + // CHECK: arith.addf + %0 = polynomial.eval #polynomial.typed_chebyshev_polynomial<[1.0, 2.0]> : !polynomial.polynomial>, %x {domain_lower = -1.0 : f64, domain_upper = 1.0 : f64} : f64 + return %0 : f64 + } +} diff --git a/tests/Transforms/polynomial_approximation/doctest.mlir b/tests/Transforms/polynomial_approximation/doctest.mlir index 9c9109e1e6..8bab7b142e 100644 --- a/tests/Transforms/polynomial_approximation/doctest.mlir +++ b/tests/Transforms/polynomial_approximation/doctest.mlir @@ -1,7 +1,7 @@ // RUN: heir-opt --polynomial-approximation %s | FileCheck %s // CHECK: @test_exp -func.func @test_exp(%x: f32) -> f32 { +func.func @test_exp(%x: f32 {secret.secret}) -> f32 { // CHECK: arith.mulf %0 = math.exp %x { degree = 3 : i32, @@ -11,9 +11,17 @@ func.func @test_exp(%x: f32) -> f32 { } // CHECK: @test_sin_default_params -func.func @test_sin_default_params(%x: f32) -> f32 { +func.func @test_sin_default_params(%x: f32 {secret.secret}) -> f32 { // CHECK: polynomial.eval // CHECK-SAME: [{{.*}}, {{.*}}, {{.*}}, {{.*}}, {{.*}}, {{.*}}] %0 = math.sin %x : f32 return %0 : f32 } + +// CHECK: @cleartext_left_alone +func.func @cleartext_left_alone(%x: f32) -> f32 { + // CHECK-NOT: polynomial.eval + // CHECK: math.sin + %0 = math.sin %x : f32 + return %0 : f32 +} diff --git a/tests/Transforms/polynomial_approximation/polynomial_approximation.mlir b/tests/Transforms/polynomial_approximation/polynomial_approximation.mlir index a2b1b2b915..6ed1148c4a 100644 --- a/tests/Transforms/polynomial_approximation/polynomial_approximation.mlir +++ b/tests/Transforms/polynomial_approximation/polynomial_approximation.mlir @@ -1,7 +1,7 @@ // RUN: heir-opt --split-input-file --polynomial-approximation %s | FileCheck %s // CHECK: @test_exp -func.func @test_exp(%x: f32) -> f32 { +func.func @test_exp(%x: f32 {secret.secret}) -> f32 { // CHECK: %[[SCALE:.*]] = arith.constant 2.500000e-01 : f32 // CHECK: %[[ONE:.*]] = arith.constant 1.000000e+00 : f32 // CHECK: %[[SCALED:.*]] = arith.mulf %{{.*}}, %[[SCALE]] : f32 @@ -16,7 +16,7 @@ func.func @test_exp(%x: f32) -> f32 { // ----- // CHECK: @test_exp_tensor -func.func @test_exp_tensor(%x: tensor<4xf32>) -> tensor<4xf32> { +func.func @test_exp_tensor(%x: tensor<4xf32> {secret.secret}) -> tensor<4xf32> { // CHECK: %[[SCALE:.*]] = arith.constant dense<7.812500e-03> : tensor<4xf32> // CHECK: %[[ONE:.*]] = arith.constant dense<1.000000e+00> : tensor<4xf32> // CHECK: %[[SCALED:.*]] = arith.mulf %{{.*}}, %[[SCALE]] : tensor<4xf32> @@ -29,7 +29,7 @@ func.func @test_exp_tensor(%x: tensor<4xf32>) -> tensor<4xf32> { // ----- // CHECK: @test_domain -func.func @test_domain(%x: f32) -> f32 { +func.func @test_domain(%x: f32 {secret.secret}) -> f32 { // CHECK: polynomial.eval // CHECK-SAME: domain_upper = 2 %0 = math.exp %x {degree = 3 : i32, domain_lower = -1.0 : f64, domain_upper = 2.0 : f64} : f32 @@ -39,7 +39,7 @@ func.func @test_domain(%x: f32) -> f32 { // ----- // CHECK: @test_sin_default_params -func.func @test_sin_default_params(%x: f32) -> f32 { +func.func @test_sin_default_params(%x: f32 {secret.secret}) -> f32 { // CHECK: polynomial.eval // CHECK-SAME: [{{.*}}, {{.*}}, {{.*}}, {{.*}}, {{.*}}, {{.*}}] %0 = math.sin %x : f32 @@ -49,7 +49,7 @@ func.func @test_sin_default_params(%x: f32) -> f32 { // ----- // CHECK: @test_maximumf -func.func @test_maximumf(%x: tensor<10xf32>) -> tensor<10xf32> { +func.func @test_maximumf(%x: tensor<10xf32> {secret.secret}) -> tensor<10xf32> { // CHECK: polynomial.eval // CHECK-NOT: arith.maximumf %c0 = arith.constant dense<0.0> : tensor<10xf32> @@ -60,7 +60,7 @@ func.func @test_maximumf(%x: tensor<10xf32>) -> tensor<10xf32> { // ----- // CHECK: @test_maximumf_domain -func.func @test_maximumf_domain(%x: tensor<10xf32>) -> tensor<10xf32> { +func.func @test_maximumf_domain(%x: tensor<10xf32> {secret.secret}) -> tensor<10xf32> { // CHECK: polynomial.eval // CHECK-SAME: domain_upper = 2 // CHECK-NOT: arith.maximumf @@ -73,7 +73,7 @@ func.func @test_maximumf_domain(%x: tensor<10xf32>) -> tensor<10xf32> { // CHECK: @test_maximumf_ignore_not_splat -func.func @test_maximumf_ignore_not_splat(%x: tensor<10xf32>) -> tensor<10xf32> { +func.func @test_maximumf_ignore_not_splat(%x: tensor<10xf32> {secret.secret}) -> tensor<10xf32> { // CHECK-NOT: polynomial.eval %c0 = arith.constant dense<[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0]> : tensor<10xf32> %0 = arith.maximumf %x, %c0 : tensor<10xf32> @@ -83,7 +83,7 @@ func.func @test_maximumf_ignore_not_splat(%x: tensor<10xf32>) -> tensor<10xf32> // ----- // CHECK: @test_maximumf_ignore_arg -func.func @test_maximumf_ignore_arg(%x: tensor<10xf32>, %y: tensor<10xf32>) -> tensor<10xf32> { +func.func @test_maximumf_ignore_arg(%x: tensor<10xf32> {secret.secret}, %y: tensor<10xf32> {secret.secret}) -> tensor<10xf32> { // CHECK-NOT: polynomial.eval %0 = arith.maximumf %x, %y : tensor<10xf32> return %0 : tensor<10xf32> @@ -92,7 +92,7 @@ func.func @test_maximumf_ignore_arg(%x: tensor<10xf32>, %y: tensor<10xf32>) -> t // ----- // CHECK: @test_log_default_params -func.func @test_log_default_params(%x: f32) -> f32 { +func.func @test_log_default_params(%x: f32 {secret.secret}) -> f32 { // CHECK: polynomial.eval // CHECK-SAME: domain_lower = 1.000000e-01 // CHECK-SAME: domain_upper = 2.000000e+00 @@ -103,7 +103,7 @@ func.func @test_log_default_params(%x: f32) -> f32 { // ----- // CHECK: @test_sqrt_default_params -func.func @test_sqrt_default_params(%x: f32) -> f32 { +func.func @test_sqrt_default_params(%x: f32 {secret.secret}) -> f32 { // CHECK: polynomial.eval // CHECK-SAME: domain_lower = 0.000000e+00 // CHECK-SAME: domain_upper = 2.000000e+00 @@ -114,7 +114,7 @@ func.func @test_sqrt_default_params(%x: f32) -> f32 { // ----- // CHECK: @test_fpowi_tensor -func.func @test_fpowi_tensor(%x: tensor<1x5xf32>) -> tensor<1x5xf32> { +func.func @test_fpowi_tensor(%x: tensor<1x5xf32> {secret.secret}) -> tensor<1x5xf32> { // CHECK: polynomial.eval // CHECK-NOT: math.fpowi %cst = arith.constant dense<2> : tensor<1x5xi64> @@ -125,7 +125,7 @@ func.func @test_fpowi_tensor(%x: tensor<1x5xf32>) -> tensor<1x5xf32> { // ----- // CHECK: @test_fpowi_scalar -func.func @test_fpowi_scalar(%x: f32) -> f32 { +func.func @test_fpowi_scalar(%x: f32 {secret.secret}) -> f32 { // CHECK: polynomial.eval // CHECK-NOT: math.fpowi %cst = arith.constant 2 : i32