forked from swiftlang/swift
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathnested_calls.sil
More file actions
103 lines (85 loc) · 4.68 KB
/
Copy pathnested_calls.sil
File metadata and controls
103 lines (85 loc) · 4.68 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
// RUN: %target-swift-frontend -emit-sil -O %s | %FileCheck %s
import Builtin
import Swift
sil_stage raw
sil @foo_prim : $@convention(thin) (Float) -> (Float, Float) {
bb0(%0 : @trivial $Float):
%1 = tuple (%0 : $Float, %0 : $Float)
return %1 : $(Float, Float)
}
sil @foo_adj : $@convention(thin) (Float, Float, Float, Float) -> Float {
bb0(%0 : @trivial $Float, %1 : @trivial $Float, %2 : @trivial $Float, %3 : @trivial $Float):
return %3 : $Float
}
sil [reverse_differentiable source 0 wrt 0 primal @foo_prim adjoint @foo_adj] @foo : $@convention(thin) (Float) -> Float {
bb0(%0 : @trivial $Float):
return %0 : $Float
}
// Nested function being called by `@func_to_diff`.
sil @nested_func_without_diffattr : $@convention(thin) (Float) -> Float {
bb0(%0 : @trivial $Float):
%1 = function_ref @foo : $@convention(thin) (Float) -> Float
%2 = apply %1(%0) : $@convention(thin) (Float) -> Float
return %2 : $Float
}
// Main function to differentiate.
sil [reverse_differentiable source 0 wrt 0] @func_to_diff : $@convention(thin) (Float) -> Float {
bb0(%0 : @trivial $Float):
%1 = function_ref @nested_func_without_diffattr : $@convention(thin) (Float) -> Float
%2 = apply %1(%0) : $@convention(thin) (Float) -> Float
%3 = function_ref @nested_func_without_diffattr : $@convention(thin) (Float) -> Float
%4 = apply %3(%2) : $@convention(thin) (Float) -> Float
%5 = tuple (%2 : $Float, %2 : $Float)
%6 = tuple_extract %5 : $(Float, Float), 0
return %6 : $Float
}
// CHECK-LABEL: struct AD__func_to_diff__Type__src_0_wrt_0 {
// CHECK-NEXT: @sil_stored var pv_0: AD__nested_func_without_diffattr__Type__src_0_wrt_0
// CHECK-NEXT: @sil_stored var v_0: Float
// CHECK-NEXT: }
// CHECK-LABEL: struct AD__nested_func_without_diffattr__Type__src_0_wrt_0 {
// CHECK-NEXT: @sil_stored var pv_0: Float
// CHECK-NEXT: @sil_stored var v_0: Float
// CHECK-NEXT: }
// CHECK-LABEL: @foo_prim : $@convention(thin) (Float) -> (Float, Float) {
// CHECK: bb0(%0 : $Float):
// CHECK: %1 = tuple (%0 : $Float, %0 : $Float)
// CHECK: return %1 : $(Float, Float)
// CHECK: }
// CHECK-LABEL: @foo_adj : $@convention(thin) (Float, Float, Float, Float) -> Float {
// CHECK: bb0(%0 : $Float, %1 : $Float, %2 : $Float, %3 : $Float):
// CHECK: return %3 : $Float
// CHECK: }
// CHECK-LABEL: [reverse_differentiable source 0 wrt 0 primal @foo_prim adjoint @foo_adj] @foo : $@convention(thin) (Float) -> Float {
// CHECK: bb0(%0 : $Float):
// CHECK: return %0 : $Float
// CHECK: }
// CHECK-LABEL: [reverse_differentiable source 0 wrt 0 primal @AD__nested_func_without_diffattr__primal_src_0_wrt_0 adjoint @AD__nested_func_without_diffattr__adjoint_src_0_wrt_0] @nested_func_without_diffattr : $@convention(thin) (Float) -> Float {
// CHECK: bb0(%0 : $Float):
// CHECK: return %0 : $Float
// CHECK: }
// CHECK-LABEL: [reverse_differentiable source 0 wrt 0 primal @AD__func_to_diff__primal_src_0_wrt_0 adjoint @AD__func_to_diff__adjoint_src_0_wrt_0] @func_to_diff : $@convention(thin) (Float) -> Float {
// CHECK: bb0(%0 : $Float):
// CHECK: return %0 : $Float
// CHECK: }
// CHECK-LABEL: @AD__func_to_diff__primal_src_0_wrt_0 : $@convention(thin) (Float) -> (@owned AD__func_to_diff__Type__src_0_wrt_0, Float) {
// CHECK: bb0(%0 : $Float):
// CHECK: %1 = struct $AD__nested_func_without_diffattr__Type__src_0_wrt_0 (%0 : $Float, %0 : $Float)
// CHECK: %2 = struct $AD__func_to_diff__Type__src_0_wrt_0 (%1 : $AD__nested_func_without_diffattr__Type__src_0_wrt_0, %0 : $Float)
// CHECK: %3 = tuple (%2 : $AD__func_to_diff__Type__src_0_wrt_0, %0 : $Float)
// CHECK: return %3 : $(AD__func_to_diff__Type__src_0_wrt_0, Float)
// CHECK: }
// CHECK-LABEL: @AD__nested_func_without_diffattr__primal_src_0_wrt_0 : $@convention(thin) (Float) -> (@owned AD__nested_func_without_diffattr__Type__src_0_wrt_0, Float) {
// CHECK: bb0(%0 : $Float):
// CHECK: %1 = struct $AD__nested_func_without_diffattr__Type__src_0_wrt_0 (%0 : $Float, %0 : $Float)
// CHECK: %2 = tuple (%1 : $AD__nested_func_without_diffattr__Type__src_0_wrt_0, %0 : $Float)
// CHECK: return %2 : $(AD__nested_func_without_diffattr__Type__src_0_wrt_0, Float)
// CHECK: }
// CHECK-LABEL: @AD__func_to_diff__adjoint_src_0_wrt_0 : $@convention(thin) (Float, AD__func_to_diff__Type__src_0_wrt_0, Float, Float) -> Float {
// CHECK: bb0(%0 : $Float, %1 : $AD__func_to_diff__Type__src_0_wrt_0, %2 : $Float, %3 : $Float):
// CHECK: return %3 : $Float
// CHECK: }
// CHECK-LABEL: @AD__nested_func_without_diffattr__adjoint_src_0_wrt_0 : $@convention(thin) (Float, AD__nested_func_without_diffattr__Type__src_0_wrt_0, Float, Float) -> Float {
// CHECK: bb0(%0 : $Float, %1 : $AD__nested_func_without_diffattr__Type__src_0_wrt_0, %2 : $Float, %3 : $Float):
// CHECK: return %3 : $Float
// CHECK: }