Skip to content

Commit a64b9f9

Browse files
fix: lower chained secret Chebyshev evaluations natively
SecretnessAnalysis runs before walk-based rewriting. Replacing an earlier eval creates a new operand with no lattice entry, while each later eval result remains analyzed. Query the result so chained secret evaluations consistently use native lowering.
1 parent ae757ed commit a64b9f9

3 files changed

Lines changed: 19 additions & 4 deletions

File tree

lib/Transforms/LowerPolynomialEval/Patterns.cpp

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -230,8 +230,8 @@ LogicalResult LowerViaPatersonStockmeyerChebyshev::matchAndRewrite(
230230

231231
LogicalResult LowerToKernelEvalChebyshev::matchAndRewrite(
232232
EvalOp op, PatternRewriter& rewriter) const {
233-
if (!mlir::heir::isSecret(op.getValue(), &solver)) {
234-
return rewriter.notifyMatchFailure(op, "operand is not secret");
233+
if (!mlir::heir::isSecret(op.getResult(), &solver)) {
234+
return rewriter.notifyMatchFailure(op, "result is not secret");
235235
}
236236
auto attr = dyn_cast<polynomial::TypedChebyshevPolynomialAttr>(
237237
op.getPolynomialAttr());

tests/Examples/lattigo/ckks/relu_composite/relu_composite_test.go

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@ import (
77
)
88

99
func TestReluComposite(t *testing.T) {
10-
evaluator, params, ecd, enc, dec := Relu_composite__configure()
10+
bootstrapEvaluator, evaluator, params, ecd, enc, dec := Relu_composite__configure()
1111

1212
const n = 16
1313
arg0 := make([]float32, n)
@@ -23,7 +23,7 @@ func TestReluComposite(t *testing.T) {
2323
ct0 := Relu_composite__encrypt__arg0(evaluator, params, ecd, enc, arg0)
2424

2525
start := time.Now()
26-
resultCt := Relu_composite(evaluator, params, ecd, ct0)
26+
resultCt := Relu_composite(bootstrapEvaluator, evaluator, params, ecd, ct0)
2727
t.Logf("composite-sign ReLU took %s", time.Since(start))
2828

2929
result := Relu_composite__decrypt__result0(evaluator, params, ecd, dec, resultCt)

tests/Transforms/lower_polynomial_eval/conditional_lower.mlir

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -18,6 +18,21 @@ module attributes {
1818
return %0 : !secret.secret<f64>
1919
}
2020

21+
// CHECK: @test_secret_kernel_chain
22+
func.func @test_secret_kernel_chain(%x: !secret.secret<f64>) -> !secret.secret<f64> {
23+
// CHECK: %[[EVAL_0:.*]] = kernel.eval_chebyshev
24+
// CHECK: %[[EVAL_1:.*]] = kernel.eval_chebyshev %[[EVAL_0]]
25+
// CHECK: kernel.eval_chebyshev %[[EVAL_1]]
26+
%0 = secret.generic(%x : !secret.secret<f64>) {
27+
^body(%x_val: f64):
28+
%1 = polynomial.eval #polynomial.typed_chebyshev_polynomial<[1.0, 2.0]> : !polynomial.polynomial<ring=<coefficientType=f64>>, %x_val {domain_lower = -1.0 : f64, domain_upper = 1.0 : f64} : f64
29+
%2 = polynomial.eval #polynomial.typed_chebyshev_polynomial<[1.0, 2.0]> : !polynomial.polynomial<ring=<coefficientType=f64>>, %1 {domain_lower = -1.0 : f64, domain_upper = 1.0 : f64} : f64
30+
%3 = polynomial.eval #polynomial.typed_chebyshev_polynomial<[1.0, 2.0]> : !polynomial.polynomial<ring=<coefficientType=f64>>, %2 {domain_lower = -1.0 : f64, domain_upper = 1.0 : f64} : f64
31+
secret.yield %3 : f64
32+
} -> (!secret.secret<f64>)
33+
return %0 : !secret.secret<f64>
34+
}
35+
2136
// CHECK: @test_public_kernel
2237
func.func @test_public_kernel(%x: f64) -> f64 {
2338
// CHECK-NOT: kernel.eval_chebyshev

0 commit comments

Comments
 (0)