Skip to content

Commit 239fe59

Browse files
authored
Cleaning & Refactoring (#508)
* remove old icnf types * more cleaning * use fitted in usage * test by more julia versions * fix * fix * more cleaning * fix maxlog * fix * fix * remove broken * cleaning * more cleaning * revert sensealg
1 parent 53da8ab commit 239fe59

23 files changed

Lines changed: 326 additions & 420 deletions

.github/workflows/CI-CheckBy.yml

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -24,8 +24,11 @@ jobs:
2424
- CheckByJET
2525
- CheckByExplicitImports
2626
version:
27-
- release
28-
- lts
27+
- "1.10"
28+
- "1.11"
29+
- "1.12"
30+
# - release
31+
# - lts
2932
# - nightly
3033
os:
3134
- ubuntu-latest

.github/workflows/CI.yml

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -27,8 +27,11 @@ jobs:
2727
- Regression
2828
- Speed
2929
version:
30-
- release
31-
- lts
30+
- "1.10"
31+
- "1.11"
32+
- "1.12"
33+
# - release
34+
# - lts
3235
# - nightly
3336
os:
3437
- ubuntu-latest

Project.toml

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -23,11 +23,12 @@ MLUtils = "f1d291b0-491e-4a28-83b9-f70985020b54"
2323
NNlib = "872c559c-99b0-510c-b3b7-b6c96a88d5cd"
2424
Optimisers = "3bd65402-5787-11e9-1adc-39752487f4e2"
2525
OptimizationOptimisers = "42dfb2eb-d2b4-4451-abcd-913932933ac1"
26-
OrdinaryDiffEqDefault = "50262376-6c5a-4cf5-baba-aaf4f84d72d7"
26+
OrdinaryDiffEqAdamsBashforthMoulton = "89bda076-bce5-4f1c-845f-551c83cdda9a"
2727
Random = "9a3f8284-a2c9-5f02-9a11-845980a1fd5c"
2828
SciMLBase = "0bca4576-84f4-4d90-8ffe-ffa030f20462"
2929
SciMLSensitivity = "1ed8b502-d754-442c-8d5d-10ac956f44a1"
3030
ScientificTypesBase = "30f210dd-8aff-4c5f-94ba-8e64358c1161"
31+
Static = "aedffcd0-7271-4cad-89d0-dc628f76c6d3"
3132
Statistics = "10745b16-79ce-11e8-11f9-7d13ad32a3b2"
3233
WeightInitializers = "d49dbf32-c5c2-4618-8acc-27bb2598ef2d"
3334
Zygote = "e88e6eb3-aa80-5325-afca-941959d7151f"
@@ -52,11 +53,12 @@ MLUtils = "0.4"
5253
NNlib = "0.9"
5354
Optimisers = "0.4"
5455
OptimizationOptimisers = "0.3"
55-
OrdinaryDiffEqDefault = "1"
56+
OrdinaryDiffEqAdamsBashforthMoulton = "1"
5657
Random = "1"
5758
SciMLBase = "2"
5859
SciMLSensitivity = "7"
5960
ScientificTypesBase = "3"
61+
Static = "1"
6062
Statistics = "1"
6163
WeightInitializers = "1"
6264
Zygote = "0.7"

benchmark/Project.toml

Lines changed: 0 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -7,9 +7,7 @@ Distributions = "31c24e10-a181-5473-b8eb-7969acd0382f"
77
ForwardDiff = "f6369f11-7733-5829-9624-2563aa707210"
88
Lux = "b2108857-7c20-44ae-9111-449ecde12c47"
99
LuxCore = "bb33d45b-7691-41d6-9220-0943567d0623"
10-
OrdinaryDiffEqDefault = "50262376-6c5a-4cf5-baba-aaf4f84d72d7"
1110
PkgBenchmark = "32113eaa-f34f-5b0d-bd6c-c81e245fc73d"
12-
SciMLSensitivity = "1ed8b502-d754-442c-8d5d-10ac956f44a1"
1311
StableRNGs = "860ef19b-820b-49d6-a774-d7a799459cd3"
1412
Zygote = "e88e6eb3-aa80-5325-afca-941959d7151f"
1513

@@ -22,9 +20,7 @@ Distributions = "0.25"
2220
ForwardDiff = "1"
2321
Lux = "1"
2422
LuxCore = "1"
25-
OrdinaryDiffEqDefault = "1"
2623
PkgBenchmark = "0.2"
27-
SciMLSensitivity = "7"
2824
StableRNGs = "1"
2925
Zygote = "0.7"
3026
julia = "1.10"

benchmark/benchmarks.jl

Lines changed: 8 additions & 31 deletions
Original file line numberDiff line numberDiff line change
@@ -6,9 +6,7 @@ import ADTypes,
66
ForwardDiff,
77
Lux,
88
LuxCore,
9-
OrdinaryDiffEqDefault,
109
PkgBenchmark,
11-
SciMLSensitivity,
1210
StableRNGs,
1311
Zygote,
1412
ContinuousNormalizingFlows
@@ -21,39 +19,18 @@ r = rand(rng, data_dist, ndimension, ndata)
2119
r = convert.(Float32, r)
2220

2321
nvars = size(r, 1)
24-
naugs = nvars
22+
naugs = nvars + 1
2523
n_in = nvars + naugs
2624

27-
nn = Lux.Chain(Lux.Dense(n_in => 3 * n_in, tanh), Lux.Dense(3 * n_in => n_in, tanh))
28-
29-
icnf = ContinuousNormalizingFlows.construct(
30-
ContinuousNormalizingFlows.ICNF,
31-
nn,
32-
nvars,
33-
naugs;
34-
compute_mode = ContinuousNormalizingFlows.LuxVecJacMatrixMode(ADTypes.AutoZygote()),
35-
tspan = (0.0f0, 1.0f0),
36-
steer_rate = 1.0f-1,
37-
λ₁ = 1.0f-2,
38-
λ₂ = 1.0f-2,
39-
λ₃ = 1.0f-2,
40-
rng,
25+
nn = Lux.Chain(
26+
Lux.Dense(n_in => (2 * n_in + 1), tanh),
27+
Lux.Dense((2 * n_in + 1) => n_in, tanh),
4128
)
4229

43-
icnf2 = ContinuousNormalizingFlows.construct(
44-
ContinuousNormalizingFlows.ICNF,
45-
nn,
46-
nvars,
47-
naugs;
48-
inplace = true,
49-
compute_mode = ContinuousNormalizingFlows.LuxVecJacMatrixMode(ADTypes.AutoZygote()),
50-
tspan = (0.0f0, 1.0f0),
51-
steer_rate = 1.0f-1,
52-
λ₁ = 1.0f-2,
53-
λ₂ = 1.0f-2,
54-
λ₃ = 1.0f-2,
55-
rng,
56-
)
30+
icnf = ContinuousNormalizingFlows.ICNF(; nn, nvars, naugmented = naugs, rng)
31+
32+
icnf2 =
33+
ContinuousNormalizingFlows.ICNF(; nn, nvars, naugmented = naugs, rng, inplace = true)
5734

5835
ps, st = LuxCore.setup(icnf.rng, icnf)
5936
ps = ComponentArrays.ComponentArray(ps)

examples/Project.toml

Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,18 @@
1+
[deps]
2+
ADTypes = "47edcb42-4c32-4615-8424-f2b9edc5f35b"
3+
CairoMakie = "13f3f980-e62b-5c42-98c6-ff1f3baf88f0"
4+
ContinuousNormalizingFlows = "00b1973d-5b2e-40bf-8604-5c9c1d8f50ac"
5+
DataFrames = "a93c6f00-e57d-5684-b7b6-d8193f3e46c0"
6+
Distances = "b4f34e82-e78d-54a5-968a-f98e89d6e8f7"
7+
Distributions = "31c24e10-a181-5473-b8eb-7969acd0382f"
8+
Logging = "56ddb016-857b-54e1-b83d-db4d58db5568"
9+
Lux = "b2108857-7c20-44ae-9111-449ecde12c47"
10+
MKL = "33e6dc65-8f57-5167-99aa-e5a354878fb2"
11+
MLDataDevices = "7e8f7934-dd98-4c1a-8fe8-92b47a384d40"
12+
MLJBase = "a7f614a8-145f-11e9-1d2a-a57a1082229d"
13+
OptimizationOptimisers = "42dfb2eb-d2b4-4451-abcd-913932933ac1"
14+
OrdinaryDiffEqAdamsBashforthMoulton = "89bda076-bce5-4f1c-845f-551c83cdda9a"
15+
SciMLSensitivity = "1ed8b502-d754-442c-8d5d-10ac956f44a1"
16+
Static = "aedffcd0-7271-4cad-89d0-dc628f76c6d3"
17+
TerminalLoggers = "5d786b92-1e48-4d6f-9151-6b4477ca9bed"
18+
Zygote = "e88e6eb3-aa80-5325-afca-941959d7151f"

examples/usage.jl

Lines changed: 46 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
# Switch To MKL For Faster Computation
2-
# using MKL
2+
using MKL
33

44
## Enable Logging
55
using Logging, TerminalLoggers
@@ -20,47 +20,63 @@ n_in = nvars + naugs
2020

2121
## Model
2222
using ContinuousNormalizingFlows,
23-
Lux, OrdinaryDiffEqAdamsBashforthMoulton, ADTypes, Zygote, MLDataDevices
23+
Lux,
24+
OrdinaryDiffEqAdamsBashforthMoulton,
25+
Static,
26+
SciMLSensitivity,
27+
ADTypes,
28+
Zygote,
29+
MLDataDevices
2430

2531
# To use gpu, add related packages
26-
# using LuxCUDA, CUDA, cuDNN
32+
# using LuxCUDA
2733

28-
nn = Chain(Dense(n_in => 3 * n_in, tanh), Dense(3 * n_in => n_in, tanh))
29-
icnf = construct(
30-
ICNF,
31-
nn,
32-
nvars, # number of variables
33-
naugs; # number of augmented dimensions
34-
compute_mode = LuxVecJacMatrixMode(AutoZygote()), # process data in batches and use Zygote
35-
inplace = false, # not using the inplace version of functions
36-
device = cpu_device(), # process data by CPU
37-
# device = gpu_device(), # process data by GPU
38-
tspan = (0.0f0, 1.0f0), # time span
39-
steer_rate = 1.0f-1, # add random noise to end of the time span
34+
nn = Chain(Dense(n_in => (2 * n_in + 1), tanh), Dense((2 * n_in + 1) => n_in, tanh))
35+
icnf = ICNF(;
36+
nn = nn,
37+
nvars = nvars, # number of variables
38+
naugmented = naugs, # number of augmented dimensions
4039
λ₁ = 1.0f-2, # regulate flow
4140
λ₂ = 1.0f-2, # regulate volume change
4241
λ₃ = 1.0f-2, # regulate augmented dimensions
43-
sol_kwargs = (; save_everystep = false, alg = VCABM()), # pass to the solver
42+
steer_rate = 1.0f-1, # add random noise to end of the time span
43+
tspan = (0.0f0, 1.0f0), # time span
44+
device = cpu_device(), # process data by CPU
45+
# device = gpu_device(), # process data by GPU
46+
cond = false, # not conditioning on auxiliary input
47+
inplace = false, # not using the inplace version of functions
48+
compute_mode = LuxVecJacMatrixMode(AutoZygote()), # process data in batches and use Zygote
49+
sol_kwargs = (;
50+
save_everystep = false,
51+
maxiters = typemax(Int),
52+
reltol = 1.0f-4,
53+
abstol = 1.0f-8,
54+
alg = VCABM(; thread = True()),
55+
sensealg = InterpolatingAdjoint(; checkpointing = true, autodiff = true),
56+
), # pass to the solver
4457
)
4558

4659
## Fit It
4760
using DataFrames, MLJBase, Zygote, ADTypes, OptimizationOptimisers
48-
df = DataFrame(transpose(r), :auto)
49-
model = ICNFModel(
50-
icnf;
51-
optimizers = (Adam(),),
52-
adtype = AutoZygote(),
53-
batchsize = 512,
54-
sol_kwargs = (; epochs = 300, progress = true), # pass to the solver
55-
)
56-
mach = machine(model, df)
57-
fit!(mach)
58-
# CUDA.@allowscalar fit!(mach) # needed for gpu
5961

60-
## Store It
6162
icnf_mach_fn = "icnf_mach.jls"
62-
MLJBase.save(icnf_mach_fn, mach) # save it
63-
mach = machine(icnf_mach_fn) # load it
63+
if ispath(icnf_mach_fn)
64+
mach = machine(icnf_mach_fn) # load it
65+
else
66+
df = DataFrame(transpose(r), :auto)
67+
model = ICNFModel(;
68+
icnf,
69+
optimizers = (OptimiserChain(WeightDecay(), Adam()),),
70+
batchsize = 1024,
71+
adtype = AutoZygote(),
72+
sol_kwargs = (; epochs = 300, progress = true), # pass to the solver
73+
)
74+
mach = machine(model, df)
75+
fit!(mach)
76+
# CUDA.@allowscalar fit!(mach) # needed for gpu
77+
78+
MLJBase.save(icnf_mach_fn, mach) # save it
79+
end
6480

6581
## Use It
6682
d = ICNFDist(mach, TestMode())

src/ContinuousNormalizingFlows.jl

Lines changed: 3 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -19,26 +19,20 @@ import ADTypes,
1919
NNlib,
2020
Optimisers,
2121
OptimizationOptimisers,
22-
OrdinaryDiffEqDefault,
22+
OrdinaryDiffEqAdamsBashforthMoulton,
2323
Random,
2424
SciMLBase,
2525
SciMLSensitivity,
2626
ScientificTypesBase,
27+
Static,
2728
Statistics,
2829
WeightInitializers,
2930
Zygote
3031

31-
export construct,
32-
inference,
32+
export inference,
3333
generate,
3434
loss,
3535
ICNF,
36-
RNODE,
37-
CondRNODE,
38-
FFJORD,
39-
CondFFJORD,
40-
Planar,
41-
CondPlanar,
4236
TestMode,
4337
TrainMode,
4438
DIVecJacVectorMode,

src/core/base_icnf.jl

Lines changed: 0 additions & 77 deletions
Original file line numberDiff line numberDiff line change
@@ -1,80 +1,3 @@
1-
function construct(
2-
aicnf::Type{<:AbstractICNF},
3-
nn::LuxCore.AbstractLuxLayer,
4-
nvars::Int,
5-
naugmented::Int = 0;
6-
data_type::Type{<:AbstractFloat} = Float32,
7-
compute_mode::ComputeMode = LuxVecJacMatrixMode(ADTypes.AutoZygote()),
8-
inplace::Bool = false,
9-
cond::Bool = aicnf <: Union{CondRNODE, CondFFJORD, CondPlanar},
10-
device::MLDataDevices.AbstractDevice = MLDataDevices.cpu_device(),
11-
basedist::Distributions.Distribution = Distributions.MvNormal(
12-
FillArrays.Zeros{data_type}(nvars + naugmented),
13-
FillArrays.Eye{data_type}(nvars + naugmented),
14-
),
15-
tspan::NTuple{2} = (zero(data_type), one(data_type)),
16-
steer_rate::AbstractFloat = zero(data_type),
17-
epsdist::Distributions.Distribution = Distributions.MvNormal(
18-
FillArrays.Zeros{data_type}(nvars + naugmented),
19-
FillArrays.Eye{data_type}(nvars + naugmented),
20-
),
21-
sol_kwargs::NamedTuple = (;),
22-
rng::Random.AbstractRNG = MLDataDevices.default_device_rng(device),
23-
λ₁::AbstractFloat = if aicnf <: Union{RNODE, CondRNODE}
24-
convert(data_type, 1.0e-2)
25-
else
26-
zero(data_type)
27-
end,
28-
λ₂::AbstractFloat = if aicnf <: Union{RNODE, CondRNODE}
29-
convert(data_type, 1.0e-2)
30-
else
31-
zero(data_type)
32-
end,
33-
λ₃::AbstractFloat = if naugmented >= nvars
34-
convert(data_type, 1.0e-2)
35-
else
36-
zero(data_type)
37-
end,
38-
)
39-
steerdist = Distributions.Uniform{data_type}(-steer_rate, steer_rate)
40-
41-
return ICNF{
42-
data_type,
43-
typeof(compute_mode),
44-
inplace,
45-
cond,
46-
!iszero(naugmented),
47-
!iszero(steer_rate),
48-
!iszero(λ₁),
49-
!iszero(λ₂),
50-
!iszero(λ₃),
51-
typeof(nn),
52-
typeof(nvars),
53-
typeof(device),
54-
typeof(basedist),
55-
typeof(tspan),
56-
typeof(steerdist),
57-
typeof(epsdist),
58-
typeof(sol_kwargs),
59-
typeof(rng),
60-
}(
61-
nn,
62-
nvars,
63-
naugmented,
64-
compute_mode,
65-
device,
66-
basedist,
67-
tspan,
68-
steerdist,
69-
epsdist,
70-
sol_kwargs,
71-
rng,
72-
λ₁,
73-
λ₂,
74-
λ₃,
75-
)
76-
end
77-
781
function Base.show(io::IO, icnf::AbstractICNF)
792
return print(io, typeof(icnf))
803
end

0 commit comments

Comments
 (0)