-
Notifications
You must be signed in to change notification settings - Fork 88
Expand file tree
/
Copy pathoptimizers.jl
More file actions
97 lines (80 loc) · 3.3 KB
/
Copy pathoptimizers.jl
File metadata and controls
97 lines (80 loc) · 3.3 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
# We use this module mostly as a placeholder for patches that should be merged into
# Optimisers.jl for Reactant compatibility.
module ReactantCompatibleOptimisers
using ..Lux: Utils
using Optimisers: Optimisers
using MLDataDevices: ReactantDevice, get_device, with_track_numbers
# We need to wrap in a "ReactantOptimiser" to correctly update the learning rate as such
# without accidentally making them constants
struct ReactantOptimiser{T} <: Optimisers.AbstractRule
opt::T
end
function Base.show(io::IO, opt::ReactantOptimiser)
print(io, "ReactantOptimiser(", opt.opt, ")")
return nothing
end
function Optimisers.apply!(opt::ReactantOptimiser, state, x, y)
return Optimisers.apply!(opt.opt, state, x, y)
end
Optimisers.init(opt::ReactantOptimiser, ps) = Optimisers.init(opt.opt, ps)
# JITing this is not great atm, since it causes issues without implementing result sharding
# annotatations
for common_opt in (:Adam, :AdaMax, :NAdam, :AdamW, :AdaBelief)
@eval function Optimisers.init(
opt::ReactantOptimiser{<:Optimisers.$(common_opt)}, x::AbstractArray{T}
) where {T}
return zero(x), zero(x), Utils.convert_eltype.((T,), opt.opt.beta)
end
end
function Optimisers.init(
opt::ReactantOptimiser{<:Optimisers.RAdam}, x::AbstractArray{T}
) where {T}
dev = with_track_numbers(get_device(opt), Integer)
return (zero(x), zero(x), Utils.convert_eltype.((T,), opt.opt.beta), dev(1))
end
function Optimisers.init(
opt::ReactantOptimiser{<:Optimisers.OAdam}, x::AbstractArray{T}
) where {T}
return zero(x), zero(x), Utils.convert_eltype.((T,), opt.opt.beta), zero(x)
end
# Recurse through chains
function Optimisers.adjust(
chain::ReactantOptimiser{<:Optimisers.OptimiserChain},
eta::Real
)
results = Optimisers.OptimiserChain([Optimisers.adjust(opt, eta) for opt in chain.opt.opts]...)
return ReactantOptimiser(results)
end
function Optimisers.adjust(
chain::ReactantOptimiser{<:Optimisers.OptimiserChain};
kw...
)
results = Optimisers.OptimiserChain([Optimisers.adjust(opt; kw...) for opt in chain.opt.opts]...)
return ReactantOptimiser(results)
end
function Optimisers._adjust(opt::ReactantOptimiser, nt::NamedTuple)
dev = with_track_numbers(get_device(opt), AbstractFloat)
return ReactantOptimiser(Optimisers._adjust(opt.opt, dev(nt)))
end
function Optimisers._adjust(opt::ReactantOptimiser{<:Optimisers.AccumGrad}, nt::NamedTuple)
dev = with_track_numbers(get_device(opt), Integer)
return ReactantOptimiser(Optimisers._adjust(opt.opt, dev(nt)))
end
# Convert existing Optimisers.jl rules to ReactantOptimisers
function make_reactant_compatible(opt::Optimisers.OptimiserChain, dev::ReactantDevice)
return ReactantOptimiser(
Optimisers.OptimiserChain(make_reactant_compatible.(opt.opts, (dev,))...)
)
end
function make_reactant_compatible(opt::Optimisers.AbstractRule, dev::ReactantDevice)
return ReactantOptimiser(with_track_numbers(dev, AbstractFloat)(opt))
end
function make_reactant_compatible(opt::Optimisers.AccumGrad, dev::ReactantDevice)
return ReactantOptimiser(with_track_numbers(dev, Integer)(opt))
end
function make_reactant_compatible(opt::Optimisers.ClipNorm, dev::ReactantDevice)
return ReactantOptimiser(
Optimisers.ClipNorm(with_track_numbers(dev, Integer)(opt.omega), opt.p, false)
)
end
end