Skip to content

Commit 17db134

Browse files
committed
fix test rng
1 parent 30ac9bc commit 17db134

3 files changed

Lines changed: 9 additions & 7 deletions

File tree

test/ci_tests/regression_tests.jl

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,8 @@
11
Test.@testset "Regression Tests" begin
2+
rng = StableRNGs.StableRNG(1)
23
ndata = 2^10
34
ndimension = 1
45
data_dist = Distributions.Beta{Float32}(2.0f0, 4.0f0)
5-
r = rand(data_dist, ndimension, ndata)
6-
r = convert.(Float32, r)
76

87
nvars = size(r, 1)
98
naugs = nvars
@@ -22,13 +21,16 @@ Test.@testset "Regression Tests" begin
2221
λ₁ = 1.0f-2,
2322
λ₂ = 1.0f-2,
2423
λ₃ = 1.0f-2,
24+
rng,
2525
sol_kwargs = (;
2626
save_everystep = false,
2727
alg = OrdinaryDiffEqDefault.DefaultODEAlgorithm(),
2828
sensealg = SciMLSensitivity.InterpolatingAdjoint(),
2929
),
3030
)
3131

32+
r = rand(icnf.rng, data_dist, ndimension, ndata)
33+
r = convert.(Float32, r)
3234
df = DataFrames.DataFrame(transpose(r), :auto)
3335
model = ContinuousNormalizingFlows.ICNFModel(
3436
icnf;

test/ci_tests/speed_tests.jl

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -21,11 +21,10 @@ Test.@testset "Speed Tests" begin
2121
Test.@testset "$compute_mode" for compute_mode in compute_modes
2222
@show compute_mode
2323

24+
rng = StableRNGs.StableRNG(1)
2425
ndata = 2^10
2526
ndimension = 1
2627
data_dist = Distributions.Beta{Float32}(2.0f0, 4.0f0)
27-
r = rand(data_dist, ndimension, ndata)
28-
r = convert.(Float32, r)
2928

3029
nvars = size(r, 1)
3130
naugs = nvars
@@ -44,15 +43,17 @@ Test.@testset "Speed Tests" begin
4443
λ₁ = 1.0f-2,
4544
λ₂ = 1.0f-2,
4645
λ₃ = 1.0f-2,
46+
rng,
4747
sol_kwargs = (;
4848
save_everystep = false,
4949
alg = OrdinaryDiffEqDefault.DefaultODEAlgorithm(),
5050
sensealg = SciMLSensitivity.InterpolatingAdjoint(),
5151
),
5252
)
5353

54+
r = rand(icnf.rng, data_dist, ndimension, ndata)
55+
r = convert.(Float32, r)
5456
df = DataFrames.DataFrame(transpose(r), :auto)
55-
5657
model = ContinuousNormalizingFlows.ICNFModel(
5758
icnf;
5859
batchsize = 0,

test/runtests.jl

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -31,8 +31,7 @@ if GROUP == "All"
3131
Logging.global_logger(debuglogger)
3232
end
3333

34-
Test.@testset verbose = true showtiming = true failfast = false rng =
35-
StableRNGs.StableRNG(1) "Overall" begin
34+
Test.@testset verbose = true showtiming = true failfast = false "Overall" begin
3635
if GROUP == "All" || GROUP in ["SmokeXOut", "SmokeXIn", "SmokeXYOut", "SmokeXYIn"]
3736
include(joinpath("ci_tests", "smoke_tests.jl"))
3837
end

0 commit comments

Comments
 (0)