Remove future steps from loss log if resuming from an earlier step
This commit is contained in:
@@ -109,6 +109,8 @@ class UILogger:
|
||||
self._flush_every_secs = float(flush_every_secs)
|
||||
self._last_flush = time.time()
|
||||
|
||||
self._first_commit_done = False
|
||||
|
||||
# start logging the training
|
||||
def start(self):
|
||||
if self._started:
|
||||
@@ -153,6 +155,12 @@ class UILogger:
|
||||
if step >= self._step_counter:
|
||||
self._step_counter = step + 1
|
||||
|
||||
# On the first commit of this run, prune any rows from a prior run
|
||||
# whose step is greater than where we are resuming from.
|
||||
if not self._first_commit_done:
|
||||
self._prune_future_steps(step)
|
||||
self._first_commit_done = True
|
||||
|
||||
wall_time = time.time()
|
||||
|
||||
# buffer step row (upsert later)
|
||||
@@ -253,6 +261,26 @@ class UILogger:
|
||||
except Exception:
|
||||
return None, str(v)
|
||||
|
||||
def _prune_future_steps(self, current_step: int) -> None:
|
||||
assert self._con is not None
|
||||
con = self._con
|
||||
|
||||
con.execute("BEGIN;")
|
||||
# metrics rows cascade via FK ON DELETE CASCADE
|
||||
con.execute("DELETE FROM steps WHERE step > ?;", (current_step,))
|
||||
# drop any keys that no longer have any metrics, and clamp last_seen_step
|
||||
con.execute(
|
||||
"DELETE FROM metric_keys "
|
||||
"WHERE NOT EXISTS (SELECT 1 FROM metrics WHERE metrics.key = metric_keys.key);"
|
||||
)
|
||||
con.execute(
|
||||
"UPDATE metric_keys "
|
||||
"SET last_seen_step = (SELECT MAX(step) FROM metrics WHERE metrics.key = metric_keys.key) "
|
||||
"WHERE last_seen_step > ?;",
|
||||
(current_step,),
|
||||
)
|
||||
con.execute("COMMIT;")
|
||||
|
||||
def _flush(self) -> None:
|
||||
if not self._pending_steps and not self._pending_metrics:
|
||||
return
|
||||
|
||||
Reference in New Issue
Block a user