Skip to content

Commit f1965e2

Browse files
committed
add lit tests
1 parent 4840031 commit f1965e2

1 file changed

Lines changed: 94 additions & 0 deletions

File tree

Lines changed: 94 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,94 @@
1+
// RUN: %eopt --enzyme --canonicalize --remove-unnecessary-enzyme-ops --enzyme-simplify-math %s | FileCheck %s
2+
3+
func.func private @scope_overwrite(%x: memref<f32>) {
4+
memref.alloca_scope {
5+
%val = memref.load %x[] : memref<f32>
6+
%cos = math.cos %val : f32
7+
memref.store %cos, %x[] : memref<f32>
8+
}
9+
return
10+
}
11+
12+
func.func @dscope_overwrite(%x: memref<f32>, %dx: memref<f32>) {
13+
enzyme.autodiff @scope_overwrite(%x, %dx) {
14+
activity = [#enzyme<activity enzyme_dup>],
15+
ret_activity = []
16+
} : (memref<f32>, memref<f32>) -> ()
17+
return
18+
}
19+
20+
func.func private @scope_load_mul(%x: memref<f64>, %c: f64) -> f64 {
21+
%r = memref.alloca_scope -> (f64) {
22+
%v = memref.load %x[] : memref<f64>
23+
%p = arith.mulf %v, %c : f64
24+
memref.alloca_scope.return %p : f64
25+
}
26+
return %r : f64
27+
}
28+
29+
func.func @dscope_load_mul(%x: memref<f64>, %dx: memref<f64>,
30+
%c: f64, %dout: f64) -> f64 {
31+
%res = enzyme.autodiff @scope_load_mul(%x, %dx, %c, %dout) {
32+
activity = [#enzyme<activity enzyme_dup>, #enzyme<activity enzyme_active>],
33+
ret_activity = [#enzyme<activity enzyme_activenoneed>]
34+
} : (memref<f64>, memref<f64>, f64, f64) -> f64
35+
return %res : f64
36+
}
37+
38+
// CHECK-LABEL: func.func private @diffescope_overwrite(
39+
// CHECK-SAME: %[[X:.*]]: memref<f32>,
40+
// CHECK-SAME: %[[DX:.*]]: memref<f32>) {
41+
// CHECK: %[[ZERO:.*]] = arith.constant 0.000000e+00 : f32
42+
// CHECK: %[[CDX0:.*]] = "enzyme.init"() : () -> !enzyme.Cache<memref<f32>>
43+
// CHECK: %[[CVAL:.*]] = "enzyme.init"() : () -> !enzyme.Cache<f32>
44+
// CHECK: %[[CDX1:.*]] = "enzyme.init"() : () -> !enzyme.Cache<memref<f32>>
45+
// CHECK: memref.alloca_scope {
46+
// CHECK: "enzyme.push"(%[[CDX0]], %[[DX]])
47+
// CHECK: %[[V:.*]] = memref.load %[[X]][] : memref<f32>
48+
// CHECK: "enzyme.push"(%[[CVAL]], %[[V]])
49+
// CHECK: %[[COS:.*]] = math.cos %[[V]] : f32
50+
// CHECK: "enzyme.push"(%[[CDX1]], %[[DX]])
51+
// CHECK: memref.store %[[COS]], %[[X]][] : memref<f32>
52+
// CHECK: }
53+
// CHECK: %[[DX_R1:.*]] = "enzyme.pop"(%[[CDX1]])
54+
// CHECK: %[[DOUT:.*]] = memref.load %[[DX_R1]][] : memref<f32>
55+
// CHECK: memref.store %[[ZERO]], %[[DX_R1]][] : memref<f32>
56+
// CHECK: %[[VAL:.*]] = "enzyme.pop"(%[[CVAL]])
57+
// CHECK: %[[SIN:.*]] = math.sin %[[VAL]] : f32
58+
// CHECK: %[[NEG:.*]] = arith.negf %[[SIN]] : f32
59+
// CHECK: %[[GRAD:.*]] = arith.mulf %[[DOUT]], %[[NEG]] : f32
60+
// CHECK: %[[DX_R0:.*]] = "enzyme.pop"(%[[CDX0]])
61+
// CHECK: %[[OLD:.*]] = memref.load %[[DX_R0]][] : memref<f32>
62+
// CHECK: %[[NEW:.*]] = arith.addf %[[OLD]], %[[GRAD]] : f32
63+
// CHECK: memref.store %[[NEW]], %[[DX_R0]][] : memref<f32>
64+
// CHECK: return
65+
66+
// CHECK-LABEL: func.func private @diffescope_load_mul(
67+
// CHECK-SAME: %[[X:[^,]+]]: memref<f64>,
68+
// CHECK-SAME: %[[DX:[^,]+]]: memref<f64>,
69+
// CHECK-SAME: %[[C:[^,]+]]: f64,
70+
// CHECK-SAME: %[[DOUT:[^)]+]]: f64) -> f64 {
71+
// CHECK: %[[CDX:.*]] = "enzyme.init"() : () -> !enzyme.Cache<memref<f64>>
72+
// CHECK: %[[GC:.*]] = "enzyme.init"() : () -> !enzyme.Gradient<f64>
73+
// CHECK: %[[CC:.*]] = "enzyme.init"() : () -> !enzyme.Cache<f64>
74+
// CHECK: %[[CV:.*]] = "enzyme.init"() : () -> !enzyme.Cache<f64>
75+
// CHECK: memref.alloca_scope -> (f64) {
76+
// CHECK: "enzyme.push"(%[[CDX]], %[[DX]])
77+
// CHECK: %[[V:.*]] = memref.load %[[X]][] : memref<f64>
78+
// CHECK: "enzyme.push"(%[[CV]], %[[V]])
79+
// CHECK: "enzyme.push"(%[[CC]], %[[C]])
80+
// CHECK: %[[P:.*]] = arith.mulf %[[V]], %[[C]] : f64
81+
// CHECK: memref.alloca_scope.return %[[P]] : f64
82+
// CHECK: }
83+
// CHECK: memref.alloca_scope {
84+
// CHECK: %[[V_R:.*]] = "enzyme.pop"(%[[CV]])
85+
// CHECK: %[[C_R:.*]] = "enzyme.pop"(%[[CC]])
86+
// CHECK-DAG: %[[DC:.*]] = arith.mulf %[[DOUT]], %[[V_R]] : f64
87+
// CHECK-DAG: %[[DV:.*]] = arith.mulf %[[DOUT]], %[[C_R]] : f64
88+
// CHECK: %[[DX_R:.*]] = "enzyme.pop"(%[[CDX]])
89+
// CHECK: %[[OLD:.*]] = memref.load %[[DX_R]][] : memref<f64>
90+
// CHECK: %[[NEW:.*]] = arith.addf %[[OLD]], %[[DV]] : f64
91+
// CHECK: memref.store %[[NEW]], %[[DX_R]][] : memref<f64>
92+
// CHECK: }
93+
// CHECK: %[[OUT:.*]] = "enzyme.get"(%[[GC]]) : (!enzyme.Gradient<f64>) -> f64
94+
// CHECK: return %[[OUT]] : f64

0 commit comments

Comments
 (0)