Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 21 additions & 16 deletions src/controls.jl
Original file line number Diff line number Diff line change
Expand Up @@ -13,15 +13,17 @@ Step(; n=5) = Step(n)
body="Train for `n` more iterations. "*
"Will never trigger a stop. ")

function update!(c::Step, model, verbosity, args...)
if verbosity > 1
@info "Stepping model for $(c.n) iterations. "
else
nothing
end
function update!(c::Step, model, verbosity, state=(n_iterations = 0,))
n_iterations = state.n_iterations
verbosity > 1 &&
@info "Stepping model for $(c.n) more iterations. "
train!(model, c.n)
state = (n_iterations = n_iterations + c.n,)
return state
end

takedown(c::Step, verbosity, state) = state

# # Info

struct Info{F<:Function}
Expand Down Expand Up @@ -286,11 +288,14 @@ WithLossDo(; f=x->@info("loss: $x"), kwargs...) = WithLossDo(f, kwargs...)

EarlyStopping.needs_loss(::Type{<:WithLossDo}) = true

function update!(c::WithLossDo, model, verbosity, state=(done=false, ))
function update!(c::WithLossDo,
model,
verbosity,
state=(loss=nothing, done=false))
loss = IterationControl.loss(model)
r = c.f(loss)
done = (c.stop_if_true && r isa Bool && r) ? true : false
return (done=done,)
return (loss=loss, done=done)
end

done(c::WithLossDo, state) = state.done
Expand All @@ -301,9 +306,9 @@ function takedown(c::WithLossDo, verbosity, state)
"Stop triggered by a `WithLossDo` control. " :
c.stop_message
verbosity > 0 && @info message
return (done = true, log = message)
return merge(state, (log = message,))
else
return (done = false, log = "")
return merge(state, (log = "",))
end
end

Expand Down Expand Up @@ -340,11 +345,11 @@ EarlyStopping.needs_training_losses(::Type{<:WithTrainingLossesDo}) = true
function update!(c::WithTrainingLossesDo,
model,
verbosity,
state=(done=false, ))
state=(latest_training_loss = nothing, done = false))
losses = IterationControl.training_losses(model)
r = c.f(losses)
done = (c.stop_if_true && r isa Bool && r) ? true : false
return (done=done, )
return (latest_training_loss=losses[end], done=done)
end

done(c::WithTrainingLossesDo, state) = state.done
Expand All @@ -355,9 +360,9 @@ function takedown(c::WithTrainingLossesDo, verbosity, state)
"Stop triggered by a `WithTrainingLossesDo` control. " :
c.stop_message
verbosity > 0 && @info message
return (done = true, log = message)
return merge(state, (log = message,))
else
return (done = false, log = "")
return merge(state, (log = "",))
end
end

Expand Down Expand Up @@ -403,8 +408,8 @@ function takedown(c::WithNumberDo, verbosity, state)
"Stop triggered by a `WithNumberDo` control. " :
c.stop_message
verbosity > 0 && @info message
return (done = true, log = message)
return merge(state, (log = message,))
else
return (done = false, log = "")
return merge(state, (log = "",))
end
end
28 changes: 18 additions & 10 deletions test/controls.jl
Original file line number Diff line number Diff line change
Expand Up @@ -6,12 +6,12 @@
m = SquareRooter(4)
c = Step(n=2)
state = IC.update!(c, m, 0)
@test state === nothing
@test state === (n_iterations = 2,)
@test m.training_losses == all_training_losses[1:2]
state = IC.update!(c, m, 0)
state = IC.update!(c, m, 0, state)
@test m.training_losses == all_training_losses[3:4]
@test !IC.done(c, state)
@test IC.takedown(c, 1, state) == NamedTuple()
@test IC.takedown(c, 1, state) == (n_iterations = 4,)
end

@testset "Info" begin
Expand Down Expand Up @@ -181,7 +181,8 @@ end
state = IC.update!(c, m, 1, state)
@test !state.done
@test v ≈ [2.25, (3281/1640)^2 - 4]
@test IC.takedown(c, 0, state) == (done = false, log="")
@test IC.takedown(c, 0, state) ==
(loss=v[end], done = false, log="")

v = Float64[]
f2(loss) = (push!(v, loss); last(v) < 0.02)
Expand All @@ -196,7 +197,8 @@ end
@test state.done
@test v ≈ [2.25, (3281/1640)^2 - 4]
@test IC.takedown(c, 0, state) ==
(done = true,
(loss = v[end],
done = true,
log="Stop triggered by a `WithLossDo` control. ")

v = Float64[]
Expand All @@ -212,7 +214,8 @@ end
@test state.done
@test v ≈ [2.25, (3281/1640)^2 - 4]
@test IC.takedown(c, 0, state) ==
(done = true,
(loss = v[end],
done = true,
log="foo")

end
Expand All @@ -231,7 +234,8 @@ end
state = IC.update!(c, m, 1, state)
@test !state.done
@test v ≈ [1.5, 0.45]
@test IC.takedown(c, 0, state) == (done = false, log="")
@test IC.takedown(c, 0, state) ==
(latest_training_loss = v[end], done = false, log="")

v = Float64[]
f1(training_loss) = (push!(v, last(training_loss)); last(v) < 0.5)
Expand All @@ -246,7 +250,8 @@ end
@test state.done
@test v ≈ [1.5, 0.45]
@test IC.takedown(c, 0, state) ==
(done = true,
(latest_training_loss = v[end],
done = true,
log="Stop triggered by a `WithTrainingLossesDo` control. ")

v = Float64[]
Expand All @@ -262,7 +267,8 @@ end
@test state.done
@test v ≈ [1.5, 0.45]
@test IC.takedown(c, 0, state) ==
(done = true,
(latest_training_loss = v[end],
done = true,
log="foo")
end

Expand All @@ -280,7 +286,7 @@ end
state = IC.update!(c, m, 1, state)
@test !state.done
@test v == [1, 2]
@test IC.takedown(c, 0, state) == (done = false, log="")
@test IC.takedown(c, 0, state) == (done = false, n = 2, log="")

v = Int[]
f2(n) = (push!(v, n); last(n) > 1)
Expand All @@ -296,6 +302,7 @@ end
@test v == [1, 2]
@test IC.takedown(c, 0, state) ==
(done = true,
n= 2,
log="Stop triggered by a `WithNumberDo` control. ")

v = Int[]
Expand All @@ -312,6 +319,7 @@ end
@test v == [1, 2]
@test IC.takedown(c, 0, state) ==
(done = true,
n = 2,
log="foo")
end

Expand Down
8 changes: 4 additions & 4 deletions test/train.jl
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
@testset "basic integration" begin
m = SquareRooter(4)
report = IC.train!(m, Step(2), InvalidValue(), NumberLimit(3); verbosity=0);
@test report[1] == (Step(2), NamedTuple())
@test report[1] == (Step(2), (n_iterations = 6,))
@test report[2] == (InvalidValue(), (done=false, log=""))
report[3] == (NumberLimit(3),
(done=true,
Expand All @@ -12,9 +12,9 @@
@test_logs((:info, r"Stop triggered by Num"),
IC.train!(m, Step(2), InvalidValue(), NumberLimit(3)));
@test_logs((:info, r"Using these controls"),
(:info, r"Stepping model for 2 iterations"),
(:info, r"Stepping model for 2 iterations"),
(:info, r"Stepping model for 2 iterations"),
(:info, r"Stepping model for 2 more iterations"),
(:info, r"Stepping model for 2 more iterations"),
(:info, r"Stepping model for 2 more iterations"),
(:info, r"Stop triggered by NumberLimit"),
IC.train!(m, Step(2),
InvalidValue(),
Expand Down