Skip to content

Commit 667b04d

Browse files
skip Mooncake for some tests
1 parent f6ccadb commit 667b04d

2 files changed

Lines changed: 18 additions & 11 deletions

File tree

lib/LuxLib/test/others/bmm/autodiff_tests.jl

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
using LuxLib, Test
22
using LuxLib: batched_matmul, batched_vec
33
using NNlib: batched_adjoint, batched_transpose
4-
using LuxTestUtils: @test_gradients, AutoEnzyme
4+
using LuxTestUtils: @test_gradients, AutoEnzyme, AutoMooncake
55

66
include("bmm_testsetup.jl")
77

@@ -30,15 +30,15 @@ include("bmm_testsetup.jl")
3030
aType(randn(rng, Float32, P, Q, B));
3131
atol=1.0e-3,
3232
rtol=1.0e-3,
33-
skip_backends=[AutoEnzyme()]
33+
skip_backends=[AutoEnzyme(), AutoMooncake()]
3434
)
3535
@test_gradients(
3636
fn,
3737
aType(randn(rng, Float32, M, P, B)),
3838
batched_transpose(aType(randn(rng, Float32, Q, P, B)));
3939
atol=1.0e-3,
4040
rtol=1.0e-3,
41-
skip_backends=[AutoEnzyme()]
41+
skip_backends=[AutoEnzyme(), AutoMooncake()]
4242
)
4343
end
4444

@@ -57,15 +57,15 @@ include("bmm_testsetup.jl")
5757
aType(randn(rng, Float32, P, Q, B));
5858
atol=1.0e-3,
5959
rtol=1.0e-3,
60-
skip_backends=[AutoEnzyme()]
60+
skip_backends=[AutoEnzyme(), AutoMooncake()]
6161
)
6262
@test_gradients(
6363
fn,
6464
aType(randn(rng, Float32, M, P)),
6565
batched_adjoint(aType(randn(rng, Float32, Q, P, B)));
6666
atol=1.0e-3,
6767
rtol=1.0e-3,
68-
skip_backends=[AutoEnzyme()]
68+
skip_backends=[AutoEnzyme(), AutoMooncake()]
6969
)
7070

7171
@test_gradients(
@@ -82,15 +82,15 @@ include("bmm_testsetup.jl")
8282
aType(randn(rng, Float32, P, Q, B));
8383
atol=1.0e-3,
8484
rtol=1.0e-3,
85-
skip_backends=[AutoEnzyme()]
85+
skip_backends=[AutoEnzyme(), AutoMooncake()]
8686
)
8787
@test_gradients(
8888
fn,
8989
aType(randn(rng, Float32, M, P)),
9090
batched_adjoint(aType(randn(rng, Float32, Q, P, B)));
9191
atol=1.0e-3,
9292
rtol=1.0e-3,
93-
skip_backends=[AutoEnzyme()]
93+
skip_backends=[AutoEnzyme(), AutoMooncake()]
9494
)
9595
end
9696

@@ -109,15 +109,15 @@ include("bmm_testsetup.jl")
109109
aType(randn(rng, Float32, P, Q, B));
110110
atol=1.0e-3,
111111
rtol=1.0e-3,
112-
skip_backends=[AutoEnzyme()]
112+
skip_backends=[AutoEnzyme(), AutoMooncake()]
113113
)
114114
@test_gradients(
115115
fn,
116116
aType(randn(rng, Float32, M, P, 1)),
117117
batched_transpose(aType(randn(rng, Float32, Q, P, B)));
118118
atol=1.0e-3,
119119
rtol=1.0e-3,
120-
skip_backends=[AutoEnzyme()]
120+
skip_backends=[AutoEnzyme(), AutoMooncake()]
121121
)
122122
end
123123
end

test/misc/helpers/loss_tests.jl

Lines changed: 9 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -156,12 +156,16 @@ end
156156
@jet celoss(ŷ, y)
157157
@jet celoss_smooth(ŷ, y)
158158

159+
# ŷ is Adjoint-wrapped; Mooncake can't rebuild that wrapper yet, so the
160+
# gradient comes back with the wrong shape.
159161
@test_gradients(
160162
Base.Fix2(celoss, y),
161163
ŷ;
162164
atol=1.0f-3,
163165
rtol=1.0f-3,
164-
skip_backends=VERSION v"1.11-" ? [AutoEnzyme()] : []
166+
skip_backends=vcat(
167+
VERSION v"1.11-" ? [AutoEnzyme()] : [], [AutoMooncake()]
168+
)
165169
)
166170
end
167171

@@ -182,12 +186,15 @@ end
182186
@jet logitceloss(logŷ, y)
183187
@jet logitceloss_smooth(logŷ, y)
184188

189+
# Same Adjoint-wrapper issue as above.
185190
@test_gradients(
186191
Base.Fix2(logitceloss, y),
187192
logŷ;
188193
atol=1.0f-3,
189194
rtol=1.0f-3,
190-
skip_backends=VERSION v"1.11-" ? [AutoEnzyme()] : []
195+
skip_backends=vcat(
196+
VERSION v"1.11-" ? [AutoEnzyme()] : [], [AutoMooncake()]
197+
)
191198
)
192199
end
193200

0 commit comments

Comments
 (0)