-
Notifications
You must be signed in to change notification settings - Fork 45
Expand file tree
/
Copy pathutils.jl
More file actions
111 lines (102 loc) · 3.49 KB
/
Copy pathutils.jl
File metadata and controls
111 lines (102 loc) · 3.49 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
104
105
106
107
108
109
110
111
# Figure out which AD backend to test
const AD = get(ENV, "AD", "All")
function test_ad(f, x, broken=(); rtol=1e-6, atol=1e-6)
for b in broken
if !(
b in (
:ForwardDiff,
:Zygote,
:Mooncake,
:ReverseDiff,
:Enzyme,
:EnzymeForward,
:EnzymeReverse,
# The `Crash` ones indicate that the error will cause a Julia crash, and
# thus we can't even run `@test_broken on it.
:EnzymeForwardCrash,
:EnzymeReverseCrash,
)
)
error("Unknown broken AD backend: $b")
end
end
finitediff = FiniteDifferences.grad(central_fdm(5, 1), f, x)[1]
if AD == "All" || AD == "ForwardDiff"
if :ForwardDiff in broken
@test_broken ForwardDiff.gradient(f, x) ≈ finitediff rtol = rtol atol = atol
else
@test ForwardDiff.gradient(f, x) ≈ finitediff rtol = rtol atol = atol
end
end
if AD == "All" || AD == "Zygote"
if :Zygote in broken
@test_broken Zygote.gradient(f, x)[1] ≈ finitediff rtol = rtol atol = atol
else
∇zygote = Zygote.gradient(f, x)[1]
@test (all(iszero, finitediff) && ∇zygote === nothing) ||
isapprox(∇zygote, finitediff; rtol=rtol, atol=atol)
end
end
if AD == "All" || AD == "ReverseDiff"
if :ReverseDiff in broken
@test_broken ReverseDiff.gradient(f, x) ≈ finitediff rtol = rtol atol = atol
else
@test ReverseDiff.gradient(f, x) ≈ finitediff rtol = rtol atol = atol
end
end
if AD == "All" || AD == "Enzyme"
forward_broken = :EnzymeForward in broken || :Enzyme in broken
reverse_broken = :EnzymeReverse in broken || :Enzyme in broken
if !(:EnzymeForwardCrash in broken)
if forward_broken
@test_broken(
Enzyme.gradient(Forward, Enzyme.Const(f), x)[1] ≈ finitediff,
rtol = rtol,
atol = atol
)
else
@test(
Enzyme.gradient(Forward, Enzyme.Const(f), x)[1] ≈ finitediff,
rtol = rtol,
atol = atol
)
end
end
if !(:EnzymeReverseCrash in broken)
if reverse_broken
@test_broken(
Enzyme.gradient(set_runtime_activity(Reverse), Enzyme.Const(f), x)[1] ≈
finitediff,
rtol = rtol,
atol = atol
)
else
@test(
Enzyme.gradient(set_runtime_activity(Reverse), Enzyme.Const(f), x)[1] ≈
finitediff,
rtol = rtol,
atol = atol
)
end
end
end
if AD == "All" || AD == "Mooncake"
rule = Mooncake.build_rrule(f, x)
if :Mooncake in broken
@test_broken isapprox(
Mooncake.value_and_gradient!!(rule, f, x)[2][2],
finitediff;
rtol=rtol,
atol=atol,
)
else
@test isapprox(
Mooncake.value_and_gradient!!(rule, f, x)[2][2],
finitediff;
rtol=rtol,
atol=atol,
)
end
end
return nothing
end