Remove future steps from loss log if resuming from an earlier step

This commit is contained in:
Jaret Burkett
2026-05-18 09:23:55 -06:00
parent e03c6e4dc9
commit b42acb988f

View File

@@ -109,6 +109,8 @@ class UILogger:
self._flush_every_secs = float(flush_every_secs) self._flush_every_secs = float(flush_every_secs)
self._last_flush = time.time() self._last_flush = time.time()
self._first_commit_done = False
# start logging the training # start logging the training
def start(self): def start(self):
if self._started: if self._started:
@@ -153,6 +155,12 @@ class UILogger:
if step >= self._step_counter: if step >= self._step_counter:
self._step_counter = step + 1 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() wall_time = time.time()
# buffer step row (upsert later) # buffer step row (upsert later)
@@ -253,6 +261,26 @@ class UILogger:
except Exception: except Exception:
return None, str(v) 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: def _flush(self) -> None:
if not self._pending_steps and not self._pending_metrics: if not self._pending_steps and not self._pending_metrics:
return return