Skip to content

Commit 6403694

Browse files
j2kuncopybara-github
authored andcommitted
dedicated pass for level analysis
PiperOrigin-RevId: 912707668
1 parent 791ed2d commit 6403694

13 files changed

Lines changed: 302 additions & 1 deletion

File tree

lib/Analysis/LevelAnalysis/LevelAnalysis.cpp

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -136,6 +136,9 @@ LogicalResult LevelAnalysis::visitOperation(
136136
};
137137

138138
LevelState resultLevel = deriveResultLevel(op, operands);
139+
if (resultLevel.isInt() && resultLevel.getInt() > levelBudget) {
140+
resultLevel = LevelState(Invalid{});
141+
}
139142
for (auto result : op->getOpResults()) {
140143
if (isa<mgmt::InitOp>(op) || isSecretInternal(op, result)) {
141144
propagate(result, resultLevel);

lib/Analysis/LevelAnalysis/LevelAnalysis.h

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -200,7 +200,9 @@ class LevelAnalysis
200200
: public dataflow::SparseForwardDataFlowAnalysis<LevelLattice>,
201201
public SecretnessAnalysisDependent<LevelAnalysis> {
202202
public:
203-
using SparseForwardDataFlowAnalysis::SparseForwardDataFlowAnalysis;
203+
LevelAnalysis(DataFlowSolver& solver, int levelBudget = 40)
204+
: dataflow::SparseForwardDataFlowAnalysis<LevelLattice>(solver),
205+
levelBudget(levelBudget) {}
204206
friend class SecretnessAnalysisDependent<LevelAnalysis>;
205207

206208
void setToEntryState(LevelLattice* lattice) override {
@@ -218,6 +220,9 @@ class LevelAnalysis
218220
void propagateIfChangedWrapper(AnalysisState* state, ChangeResult changed) {
219221
propagateIfChanged(state, changed);
220222
}
223+
224+
private:
225+
int levelBudget;
221226
};
222227

223228
LevelState deriveResultLevel(Operation* op,
Lines changed: 66 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,66 @@
1+
#include "lib/Transforms/AnnotateLevel/AnnotateLevel.h"
2+
3+
#include "lib/Analysis/LevelAnalysis/LevelAnalysis.h"
4+
#include "lib/Analysis/SecretnessAnalysis/SecretnessAnalysis.h"
5+
#include "lib/Utils/AttributeUtils.h"
6+
#include "lib/Utils/Utils.h"
7+
#include "llvm/include/llvm/Support/Debug.h" // from @llvm-project
8+
#include "mlir/include/mlir/Analysis/DataFlow/Utils.h" // from @llvm-project
9+
#include "mlir/include/mlir/Analysis/DataFlowFramework.h" // from @llvm-project
10+
#include "mlir/include/mlir/IR/Builders.h" // from @llvm-project
11+
#include "mlir/include/mlir/IR/BuiltinAttributes.h" // from @llvm-project
12+
#include "mlir/include/mlir/IR/MLIRContext.h" // from @llvm-project
13+
#include "mlir/include/mlir/IR/Value.h" // from @llvm-project
14+
#include "mlir/include/mlir/Support/LLVM.h" // from @llvm-project
15+
16+
#define DEBUG_TYPE "annotate-level"
17+
18+
namespace mlir {
19+
namespace heir {
20+
21+
#define GEN_PASS_DEF_ANNOTATELEVEL
22+
#include "lib/Transforms/AnnotateLevel/AnnotateLevel.h.inc"
23+
24+
struct AnnotateLevel : impl::AnnotateLevelBase<AnnotateLevel> {
25+
using AnnotateLevelBase::AnnotateLevelBase;
26+
27+
void runOnOperation() override {
28+
DataFlowSolver solver;
29+
dataflow::loadBaselineAnalyses(solver);
30+
solver.load<SecretnessAnalysis>();
31+
solver.load<LevelAnalysis>(levelBudget);
32+
33+
auto result = solver.initializeAndRun(getOperation());
34+
35+
if (failed(result)) {
36+
getOperation()->emitOpError() << "Failed to run the analysis.\n";
37+
signalPassFailure();
38+
return;
39+
}
40+
41+
walkValues(getOperation(), [&](Value value) {
42+
auto* lattice = solver.lookupState<LevelLattice>(value);
43+
if (!lattice) return;
44+
45+
auto& state = lattice->getValue();
46+
if (!state.isInitialized()) return;
47+
48+
OpBuilder b(value.getContext());
49+
Attribute attr;
50+
if (state.isInvalid()) {
51+
attr = b.getStringAttr("invalid");
52+
} else if (state.isMaxLevel()) {
53+
attr = b.getStringAttr("max");
54+
} else if (state.isInt()) {
55+
attr = b.getIndexAttr(state.getInt());
56+
}
57+
58+
if (attr) {
59+
setAttributeAssociatedWith(value, "mgmt.level", attr);
60+
}
61+
});
62+
}
63+
};
64+
65+
} // namespace heir
66+
} // namespace mlir
Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,18 @@
1+
#ifndef LIB_TRANSFORMS_ANNOTATELEVEL_ANNOTATELEVEL_H_
2+
#define LIB_TRANSFORMS_ANNOTATELEVEL_ANNOTATELEVEL_H_
3+
4+
#include "mlir/include/mlir/Pass/Pass.h" // from @llvm-project
5+
6+
namespace mlir {
7+
namespace heir {
8+
9+
#define GEN_PASS_DECL
10+
#include "lib/Transforms/AnnotateLevel/AnnotateLevel.h.inc"
11+
12+
#define GEN_PASS_REGISTRATION
13+
#include "lib/Transforms/AnnotateLevel/AnnotateLevel.h.inc"
14+
15+
} // namespace heir
16+
} // namespace mlir
17+
18+
#endif // LIB_TRANSFORMS_ANNOTATELEVEL_ANNOTATELEVEL_H_
Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,26 @@
1+
#ifndef LIB_TRANSFORMS_ANNOTATELEVEL_ANNOTATELEVEL_TD_
2+
#define LIB_TRANSFORMS_ANNOTATELEVEL_ANNOTATELEVEL_TD_
3+
4+
include "mlir/Pass/PassBase.td"
5+
6+
def AnnotateLevel : Pass<"annotate-level"> {
7+
let summary = "Annotate level lattice values in the IR";
8+
let description = [{
9+
This pass that runs the level analysis and annotates the IR with the
10+
results, attaching `{mgmt.level}` attributes to each SSA value with
11+
an inferred value.
12+
13+
If the `level-budget` is included, it gives an upper limit on the
14+
number of iterations the level analysis will attempt before giving
15+
up and reporting a default (top) value, annotated in the resulting
16+
IR as "invalid."
17+
18+
(* example filepath=tests/Transforms/annotate_level/doctest.mlir *)
19+
}];
20+
let options = [
21+
Option<"levelBudget", "level-budget", "int", /*default=*/"40",
22+
"The level budget bound for the analysis.">,
23+
];
24+
}
25+
26+
#endif // LIB_TRANSFORMS_ANNOTATELEVEL_ANNOTATELEVEL_TD_

lib/Transforms/AnnotateLevel/BUILD

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,32 @@
1+
load("@heir//lib/Transforms:transforms.bzl", "add_heir_transforms")
2+
load("@rules_cc//cc:cc_library.bzl", "cc_library")
3+
4+
package(
5+
default_applicable_licenses = ["@heir//:license"],
6+
default_visibility = ["//visibility:public"],
7+
)
8+
9+
cc_library(
10+
name = "AnnotateLevel",
11+
srcs = ["AnnotateLevel.cpp"],
12+
hdrs = [
13+
"AnnotateLevel.h",
14+
],
15+
deps = [
16+
":pass_inc_gen",
17+
"@heir//lib/Analysis/LevelAnalysis",
18+
"@heir//lib/Analysis/SecretnessAnalysis",
19+
"@heir//lib/Utils",
20+
"@heir//lib/Utils:AttributeUtils",
21+
"@llvm-project//llvm:Support",
22+
"@llvm-project//mlir:Analysis",
23+
"@llvm-project//mlir:IR",
24+
"@llvm-project//mlir:Pass",
25+
"@llvm-project//mlir:Support",
26+
],
27+
)
28+
29+
add_heir_transforms(
30+
generated_target_name = "pass_inc_gen",
31+
pass_name = "AnnotateLevel",
32+
)
Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,10 @@
1+
load("//bazel:lit.bzl", "glob_lit_tests")
2+
3+
package(default_applicable_licenses = ["@heir//:license"])
4+
5+
glob_lit_tests(
6+
name = "all_tests",
7+
data = ["@heir//tests:test_utilities"],
8+
driver = "@heir//tests:run_lit.sh",
9+
test_file_exts = ["mlir"],
10+
)
Lines changed: 70 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,70 @@
1+
// RUN: heir-opt --annotate-level="level-budget=5" %s | FileCheck %s --check-prefix=CHECK-B5
2+
3+
module {
4+
// Part 1: No loop, at least three levels
5+
func.func @no_loop(%arg0: !secret.secret<tensor<16xf32>>) -> !secret.secret<tensor<16xf32>> {
6+
%0 = secret.generic(%arg0 : !secret.secret<tensor<16xf32>>) {
7+
^body(%val: tensor<16xf32>):
8+
// CHECK-B5: mgmt.modreduce
9+
// CHECK-B5-SAME: {mgmt.level = 1 : index}
10+
%1 = mgmt.modreduce %val : tensor<16xf32>
11+
12+
// CHECK-B5: mgmt.modreduce
13+
// CHECK-B5-SAME: {mgmt.level = 2 : index}
14+
%2 = mgmt.modreduce %1 : tensor<16xf32>
15+
16+
// CHECK-B5: mgmt.modreduce
17+
// CHECK-B5-SAME: {mgmt.level = 3 : index}
18+
%3 = mgmt.modreduce %2 : tensor<16xf32>
19+
20+
secret.yield %3 : tensor<16xf32>
21+
} -> !secret.secret<tensor<16xf32>>
22+
return %0 : !secret.secret<tensor<16xf32>>
23+
}
24+
25+
// Part 2: Loop that converges
26+
func.func @loop_converges(%arg0: !secret.secret<tensor<16xf32>>) -> !secret.secret<tensor<16xf32>> {
27+
%cst = arith.constant dense<0.000000e+00> : tensor<16xf32>
28+
%c0 = arith.constant 0 : index
29+
%c10 = arith.constant 10 : index
30+
%c1 = arith.constant 1 : index
31+
%cst_0 = arith.constant dense<1.100000e+00> : tensor<16xf32>
32+
33+
%0 = secret.generic(%arg0 : !secret.secret<tensor<16xf32>>) {
34+
^body(%val: tensor<16xf32>):
35+
%1 = scf.for %arg1 = %c0 to %c10 step %c1 iter_args(%arg2 = %cst) -> (tensor<16xf32>) {
36+
%2 = arith.mulf %val, %cst_0 : tensor<16xf32>
37+
// CHECK-B5: mgmt.modreduce
38+
// CHECK-B5-SAME: {mgmt.level = 1 : index}
39+
%3 = mgmt.modreduce %2 : tensor<16xf32>
40+
41+
// CHECK-B5: arith.addf
42+
// CHECK-B5-SAME: {mgmt.level = 1 : index}
43+
%4 = arith.addf %arg2, %3 : tensor<16xf32>
44+
scf.yield %4 : tensor<16xf32>
45+
}
46+
secret.yield %1 : tensor<16xf32>
47+
} -> !secret.secret<tensor<16xf32>>
48+
return %0 : !secret.secret<tensor<16xf32>>
49+
}
50+
51+
// Part 3: Loop that does not converge (grows to 10)
52+
func.func @loop_diverges(%arg0: !secret.secret<tensor<16xf32>>) -> !secret.secret<tensor<16xf32>> {
53+
%c0 = arith.constant 0 : index
54+
%c10 = arith.constant 10 : index
55+
%c1 = arith.constant 1 : index
56+
57+
%0 = secret.generic(%arg0 : !secret.secret<tensor<16xf32>>) {
58+
^body(%val: tensor<16xf32>):
59+
// CHECK-B5: scf.for
60+
%1 = scf.for %arg1 = %c0 to %c10 step %c1 iter_args(%arg2 = %val) -> (tensor<16xf32>) {
61+
// CHECK-B5: mgmt.modreduce
62+
// CHECK-B5-SAME: {mgmt.level = "invalid"}
63+
%2 = mgmt.modreduce %arg2 : tensor<16xf32>
64+
scf.yield %2 : tensor<16xf32>
65+
}
66+
secret.yield %1 : tensor<16xf32>
67+
} -> !secret.secret<tensor<16xf32>>
68+
return %0 : !secret.secret<tensor<16xf32>>
69+
}
70+
}
Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,23 @@
1+
// RUN: heir-opt --annotate-level="level-budget=2" %s | FileCheck %s
2+
3+
module {
4+
func.func @level_growth(%arg0: !secret.secret<tensor<16xf32>>) -> !secret.secret<tensor<16xf32>> {
5+
%cst = arith.constant dense<0.000000e+00> : tensor<16xf32>
6+
%c0 = arith.constant 0 : index
7+
%c10 = arith.constant 10 : index
8+
%c1 = arith.constant 1 : index
9+
10+
%0 = secret.generic(%arg0 : !secret.secret<tensor<16xf32>>) {
11+
^body(%val: tensor<16xf32>):
12+
// CHECK: scf.for
13+
%1 = scf.for %arg1 = %c0 to %c10 step %c1 iter_args(%arg2 = %val) -> (tensor<16xf32>) {
14+
// CHECK: mgmt.modreduce
15+
// CHECK-SAME: {mgmt.level = "invalid"}
16+
%2 = mgmt.modreduce %arg2 : tensor<16xf32>
17+
scf.yield %2 : tensor<16xf32>
18+
}
19+
secret.yield %1 : tensor<16xf32>
20+
} -> !secret.secret<tensor<16xf32>>
21+
return %0 : !secret.secret<tensor<16xf32>>
22+
}
23+
}
Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,30 @@
1+
// RUN: heir-opt --annotate-level %s | FileCheck %s
2+
3+
module {
4+
func.func @loop_matvec(%arg0: !secret.secret<tensor<16xf32>>) -> !secret.secret<tensor<16xf32>> {
5+
%cst = arith.constant dense<0.000000e+00> : tensor<16xf32>
6+
%c0 = arith.constant 0 : index
7+
%c10 = arith.constant 10 : index
8+
%c1 = arith.constant 1 : index
9+
%cst_0 = arith.constant dense<1.100000e+00> : tensor<16xf32>
10+
11+
%0 = secret.generic(%arg0 : !secret.secret<tensor<16xf32>>) {
12+
^body(%val: tensor<16xf32>):
13+
// CHECK: scf.for
14+
%1 = scf.for %arg1 = %c0 to %c10 step %c1 iter_args(%arg2 = %cst) -> (tensor<16xf32>) {
15+
// CHECK: arith.mulf
16+
// CHECK-SAME: {mgmt.level = 0 : index}
17+
%2 = arith.mulf %val, %cst_0 : tensor<16xf32>
18+
// CHECK: mgmt.modreduce
19+
// CHECK-SAME: {mgmt.level = 1 : index}
20+
%3 = mgmt.modreduce %2 : tensor<16xf32>
21+
// CHECK: arith.addf
22+
// CHECK-SAME: {mgmt.level = 1 : index}
23+
%4 = arith.addf %arg2, %3 : tensor<16xf32>
24+
scf.yield %4 : tensor<16xf32>
25+
}
26+
secret.yield %1 : tensor<16xf32>
27+
} -> !secret.secret<tensor<16xf32>>
28+
return %0 : !secret.secret<tensor<16xf32>>
29+
}
30+
}

0 commit comments

Comments
 (0)