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
2 changes: 1 addition & 1 deletion Project.toml
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
name = "IterationControl"
uuid = "b3c1a2ee-3fec-4384-bf48-272ea71de57c"
authors = ["Anthony D. Blaom <anthony.blaom@gmail.com>"]
version = "0.2.2"
version = "0.3.0"

[deps]
EarlyStopping = "792122b4-ca99-40de-a6bc-6742525f08b6"
Expand Down
4 changes: 2 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -233,12 +233,12 @@ function train!(model, controls...; verbosity::Int=1)
verbosity > 1 && @info "Using these controls: $(flat(control)). "

# first training event:
state = update!(control, model, verbosity - 1)
state = update!(control, model, verbosity)
finished = done(control, state)

# subsequent training events:
while !finished
state = update!(control, model, verbosity - 1, state)
state = update!(control, model, verbosity, state)
finished = done(control, state)
end

Expand Down
12 changes: 6 additions & 6 deletions src/controls.jl
Original file line number Diff line number Diff line change
Expand Up @@ -14,7 +14,7 @@ Step(; n=5) = Step(n)
"Will never trigger a stop. ")

function update!(c::Step, model, verbosity, args...)
if verbosity > 0
if verbosity > 1
@info "Steping model for $(c.n) iterations. "
else
nothing
Expand Down Expand Up @@ -44,7 +44,7 @@ Info(; f::Function=identity) = Info(f)
"See also [`Warn`](@ref), [`Error`](@ref). ")

function update!(c::Info, model, verbosity, args...)
verbosity < 0 || @info _log_eval(c.f, model)
verbosity < 1 || @info _log_eval(c.f, model)
return nothing
end

Expand Down Expand Up @@ -73,15 +73,15 @@ Warn(predicate; f="") = Warn(predicate, f)
"See also [`Info`](@ref), [`Error`](@ref). ")

function update!(c::Warn, model, verbosity, args...)
verbosity > 0 && c.predicate(model) &&
verbosity > 1 && c.predicate(model) &&
@warn _log_eval(c.f, model)
return nothing
end

function update!(c::Warn, model, verbosity, warnings=())
if c.predicate(model)
warning = _log_eval(c.f, model)
verbosity < 0 || @warn warning
verbosity < 1 || @warn warning
state = tuple(warnings..., warning)
else
state = warnings
Expand Down Expand Up @@ -220,7 +220,7 @@ function update!(c::Data, model, verbosity)
data_exhausted = false
item, iter_state = next
end
data_exhausted && verbosity > -1 && !c.stop_when_exhausted &&
data_exhausted && verbosity > 0 && !c.stop_when_exhausted &&
@info DATA_EXHAUSTED
data_exhausted || ingest!(model, item)
done = data_exhausted && c.stop_when_exhausted
Expand All @@ -240,7 +240,7 @@ function update!(c::Data, model, verbosity, state)
item, iter_state = next
end
end
data_exhausted && verbosity > -1 && !c.stop_when_exhausted &&
data_exhausted && verbosity > 0 && !c.stop_when_exhausted &&
iter_state !== nothing && @info DATA_EXHAUSTED
data_exhausted || ingest!(model, item)
done = data_exhausted && c.stop_when_exhausted
Expand Down
4 changes: 2 additions & 2 deletions src/train.jl
Original file line number Diff line number Diff line change
Expand Up @@ -11,12 +11,12 @@ function train!(model, controls...; verbosity::Int=1)
verbosity > 1 && @info "Using these controls: $(flat(control)). "

# first training event:
state = update!(control, model, verbosity - 1)
state = update!(control, model, verbosity)
finished = done(control, state)

# subsequent training events:
while !finished
state = update!(control, model, verbosity - 1, state)
state = update!(control, model, verbosity, state)
finished = done(control, state)
end

Expand Down
92 changes: 46 additions & 46 deletions test/controls.jl
Original file line number Diff line number Diff line change
Expand Up @@ -18,11 +18,11 @@ end
m = SquareRooter(4)
c = Info(m->m.root)
IC.train!(m, 1)
@test_logs (:info, 2.5) IC.update!(c, m, 2)
@test_logs (:info, 2.5) IC.update!(c, m, 1)
@test_logs (:info, 2.5) IC.update!(c, m, 0)
state = @test_logs IC.update!(c, m, -1)
state = @test_logs IC.update!(c, m, 0)
@test state === nothing
@test_logs (:info, 2.5) IC.update!(c, m, 1, state)
@test_logs (:info, 2.5) IC.update!(c, m, 2, state)
@test !IC.done(c, state)
@test IC.takedown(c, 10, state) == NamedTuple()
end
Expand All @@ -32,27 +32,27 @@ end
c = Warn(m -> m.root > 2.4)

IC.train!(m, 1)
@test_logs (:warn, "") IC.update!(c, m, 2)
@test_logs (:warn, "") IC.update!(c, m, 1)
@test_logs (:warn, "") IC.update!(c, m, 0)
state = @test_logs IC.update!(c, m, -1)
state = @test_logs IC.update!(c, m, 0)
@test state === ("", )

IC.train!(m, 1)
@test_logs IC.update!(c, m, 2)
@test_logs IC.update!(c, m, 1)
@test_logs IC.update!(c, m, 0)
state = @test_logs IC.update!(c, m, -1)
state = @test_logs IC.update!(c, m, 0)
@test state === ()

m = SquareRooter(4)
IC.train!(m, 1)
state = IC.update!(c, m, -1)
@test_logs (:warn, "") IC.update!(c, m, 2, state)
@test_logs (:warn, "") IC.update!(c, m, 1, state)
@test_logs (:warn, "") IC.update!(c, m, 0, state)
state = @test_logs IC.update!(c, m, -1, state)
state = @test_logs IC.update!(c, m, 0, state)
@test state === ("", "")

IC.train!(m, 1)
@test_logs IC.update!(c, m, 1, state)
@test_logs IC.update!(c, m, 2, state)
@test_logs IC.update!(c, m, 0, state)
state = @test_logs IC.update!(c, m, -, state)
@test state === ("", "")
Expand All @@ -61,14 +61,14 @@ end
c = Warn(m -> m.root > 2.4, f = m->m.root)

IC.train!(m, 1)
@test_logs (:warn, 2.5) IC.update!(c, m, 2)
@test_logs (:warn, 2.5) IC.update!(c, m, 1)
@test_logs (:warn, 2.5) IC.update!(c, m, 0)
state = @test_logs IC.update!(c, m, -1)
state = @test_logs IC.update!(c, m, 0)
@test state === (2.5, )

@test_logs (:warn, 2.5) IC.update!(c, m, 2, state)
@test_logs (:warn, 2.5) IC.update!(c, m, 1, state)
@test_logs (:warn, 2.5) IC.update!(c, m, 0, state)
state = @test_logs IC.update!(c, m, -1, state)
state = @test_logs IC.update!(c, m, 0, state)
@test state === (2.5, 2.5)

@test !IC.done(c, state)
Expand All @@ -80,34 +80,34 @@ end
c = Error(m -> m.root > 2.4)

IC.train!(m, 1)
state = @test_logs (:error, "") IC.update!(c, m, 1)
state = @test_logs (:error, "") IC.update!(c, m, 2)
@test state === (done=true, error="")

IC.train!(m, 1)
state = @test_logs IC.update!(c, m, 1)
state = @test_logs IC.update!(c, m, 2)
@test state === (done=false, error=())

m = SquareRooter(4)
IC.train!(m, 1)
state = @test_logs (:error, "") IC.update!(c, m, 1)
state = @test_logs (:error, "") IC.update!(c, m, 1, state)
state = @test_logs (:error, "") IC.update!(c, m, 2)
state = @test_logs (:error, "") IC.update!(c, m, 2, state)
@test state === (done=true, error="")

m = SquareRooter(4)
c = Error(m -> m.root > 2.4, f = m->m.root)

IC.train!(m, 1)
state = @test_logs (:error, 2.5) IC.update!(c, m, 1)
state = @test_logs (:error, 2.5) IC.update!(c, m, 2)
@test state === (done=true, error=2.5)

IC.train!(m, 1)
state = @test_logs IC.update!(c, m, 1)
state = @test_logs IC.update!(c, m, 2)
@test state === (done=false, error=())

m = SquareRooter(4)
IC.train!(m, 1)
state = @test_logs (:error, 2.5) IC.update!(c, m, 1)
state = @test_logs (:error, 2.5) IC.update!(c, m, 1, state)
state = @test_logs (:error, 2.5) IC.update!(c, m, 2)
state = @test_logs (:error, 2.5) IC.update!(c, m, 2, state)
@test state === (done=true, error=2.5)

@test IC.done(c, state)
Expand All @@ -122,11 +122,11 @@ end
c = Callback(f)
m = SquareRooter(4)
IC.train!(m, 1)
state = IC.update!(c, m, 0)
state = IC.update!(c, m, 1)
@test !state.done
@test v == [2.25, ]
IC.train!(m, 2)
state = IC.update!(c, m, 0, state)
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="")
Expand All @@ -137,11 +137,11 @@ end
c = Callback(f, stop_if_true=true)
m = SquareRooter(4)
IC.train!(m, 1)
state = IC.update!(c, m, 0)
state = IC.update!(c, m, 1)
@test !state.done
@test v == [2.25, ]
IC.train!(m, 2)
state = IC.update!(c, m, 0, state)
state = IC.update!(c, m, 1, state)
@test state.done
@test v ≈ [2.25, (3281/1640)^2 - 4]
@test IC.takedown(c, 0, state) ==
Expand All @@ -154,11 +154,11 @@ end
c = Callback(f, stop_if_true=true, stop_message="foo")
m = SquareRooter(4)
IC.train!(m, 1)
state = IC.update!(c, m, 0)
state = IC.update!(c, m, 1)
@test !state.done
@test v == [2.25, ]
IC.train!(m, 2)
state = IC.update!(c, m, 0, state)
state = IC.update!(c, m, 1, state)
@test state.done
@test v ≈ [2.25, (3281/1640)^2 - 4]
@test IC.takedown(c, 0, state) ==
Expand All @@ -174,11 +174,11 @@ end
c = WithLossDo(f)
m = SquareRooter(4)
IC.train!(m, 1)
state = IC.update!(c, m, 0)
state = IC.update!(c, m, 1)
@test !state.done
@test v == [2.25, ]
IC.train!(m, 2)
state = IC.update!(c, m, 0, state)
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="")
Expand All @@ -188,11 +188,11 @@ end
c = WithLossDo(f, stop_if_true=true)
m = SquareRooter(4)
IC.train!(m, 1)
state = IC.update!(c, m, 0)
state = IC.update!(c, m, 1)
@test !state.done
@test v == [2.25, ]
IC.train!(m, 2)
state = IC.update!(c, m, 0, state)
state = IC.update!(c, m, 1, state)
@test state.done
@test v ≈ [2.25, (3281/1640)^2 - 4]
@test IC.takedown(c, 0, state) ==
Expand All @@ -204,11 +204,11 @@ end
c = WithLossDo(f, stop_if_true=true, stop_message="foo")
m = SquareRooter(4)
IC.train!(m, 1)
state = IC.update!(c, m, 0)
state = IC.update!(c, m, 1)
@test !state.done
@test v == [2.25, ]
IC.train!(m, 2)
state = IC.update!(c, m, 0, state)
state = IC.update!(c, m, 1, state)
@test state.done
@test v ≈ [2.25, (3281/1640)^2 - 4]
@test IC.takedown(c, 0, state) ==
Expand All @@ -224,11 +224,11 @@ end
c = WithTrainingLossesDo(f)
m = SquareRooter(4)
IC.train!(m, 1)
state = IC.update!(c, m, 0)
state = IC.update!(c, m, 1)
@test !state.done
@test v ≈ [1.5, ]
IC.train!(m, 1)
state = IC.update!(c, m, 0, state)
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="")
Expand All @@ -238,11 +238,11 @@ end
c = WithTrainingLossesDo(f, stop_if_true=true)
m = SquareRooter(4)
IC.train!(m, 1)
state = IC.update!(c, m, 0)
state = IC.update!(c, m, 1)
@test !state.done
@test v == [1.5, ]
IC.train!(m, 1)
state = IC.update!(c, m, 0, state)
state = IC.update!(c, m, 1, state)
@test state.done
@test v ≈ [1.5, 0.45]
@test IC.takedown(c, 0, state) ==
Expand All @@ -254,11 +254,11 @@ end
c = WithTrainingLossesDo(f, stop_if_true=true, stop_message="foo")
m = SquareRooter(4)
IC.train!(m, 1)
state = IC.update!(c, m, 0)
state = IC.update!(c, m, 1)
@test !state.done
@test v == [1.5, ]
IC.train!(m, 1)
state = IC.update!(c, m, 0, state)
state = IC.update!(c, m, 1, state)
@test state.done
@test v ≈ [1.5, 0.45]
@test IC.takedown(c, 0, state) ==
Expand All @@ -273,11 +273,11 @@ end
c = WithNumberDo(f)
m = SquareRooter(4)
IC.train!(m, 1)
state = IC.update!(c, m, 0)
state = IC.update!(c, m, 1)
@test !state.done
@test v == [1, ]
IC.train!(m, 1)
state = IC.update!(c, m, 0, state)
state = IC.update!(c, m, 1, state)
@test !state.done
@test v == [1, 2]
@test IC.takedown(c, 0, state) == (done = false, log="")
Expand All @@ -287,11 +287,11 @@ end
c = WithNumberDo(f, stop_if_true=true)
m = SquareRooter(4)
IC.train!(m, 1)
state = IC.update!(c, m, 0)
state = IC.update!(c, m, 1)
@test !state.done
@test v == [1, ]
IC.train!(m, 1)
state = IC.update!(c, m, 0, state)
state = IC.update!(c, m, 1, state)
@test state.done
@test v == [1, 2]
@test IC.takedown(c, 0, state) ==
Expand All @@ -303,11 +303,11 @@ end
c = WithNumberDo(f, stop_if_true=true, stop_message="foo")
m = SquareRooter(4)
IC.train!(m, 1)
state = IC.update!(c, m, 0)
state = IC.update!(c, m, 1)
@test !state.done
@test v == [1, ]
IC.train!(m, 1)
state = IC.update!(c, m, 0, state)
state = IC.update!(c, m, 1, state)
@test state.done
@test v == [1, 2]
@test IC.takedown(c, 0, state) ==
Expand Down
Loading