Reza2kn commited on
Commit
3456ea9
·
verified ·
1 Parent(s): d03de05

Add checkpoint uploads to eval monitor

Browse files
Files changed (1) hide show
  1. training-kit/monitor_ppocr_eval.py +32 -1
training-kit/monitor_ppocr_eval.py CHANGED
@@ -7,6 +7,7 @@ import argparse
7
  import os
8
  import re
9
  import signal
 
10
  import time
11
  from pathlib import Path
12
 
@@ -24,6 +25,11 @@ def main() -> None:
24
  parser.add_argument("--divergence-delta", type=float, default=0.03)
25
  parser.add_argument("--divergence-patience", type=int, default=2)
26
  parser.add_argument("--initial-best", type=float, default=-1.0)
 
 
 
 
 
27
  args = parser.parse_args()
28
 
29
  best = args.initial_best
@@ -41,7 +47,8 @@ def main() -> None:
41
  scores = [float(value) for value in PATTERN.findall(args.log.read_text(errors="ignore"))]
42
  for score in scores[seen:]:
43
  seen += 1
44
- if score > best + args.min_delta:
 
45
  best, stale = score, 0
46
  else:
47
  stale += 1
@@ -51,6 +58,30 @@ def main() -> None:
51
  f"stale={stale} diverged={diverged}",
52
  flush=True,
53
  )
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
54
  if seen >= args.min_evals and (
55
  stale >= args.patience or diverged >= args.divergence_patience
56
  ):
 
7
  import os
8
  import re
9
  import signal
10
+ import subprocess
11
  import time
12
  from pathlib import Path
13
 
 
25
  parser.add_argument("--divergence-delta", type=float, default=0.03)
26
  parser.add_argument("--divergence-patience", type=int, default=2)
27
  parser.add_argument("--initial-best", type=float, default=-1.0)
28
+ parser.add_argument("--repo-id")
29
+ parser.add_argument("--checkpoint-prefix", type=Path)
30
+ parser.add_argument("--config", type=Path)
31
+ parser.add_argument("--dictionary", type=Path)
32
+ parser.add_argument("--uploader", type=Path)
33
  args = parser.parse_args()
34
 
35
  best = args.initial_best
 
47
  scores = [float(value) for value in PATTERN.findall(args.log.read_text(errors="ignore"))]
48
  for score in scores[seen:]:
49
  seen += 1
50
+ improved = score > best + args.min_delta
51
+ if improved:
52
  best, stale = score, 0
53
  else:
54
  stale += 1
 
58
  f"stale={stale} diverged={diverged}",
59
  flush=True,
60
  )
61
+ upload_values = (
62
+ args.repo_id,
63
+ args.checkpoint_prefix,
64
+ args.config,
65
+ args.dictionary,
66
+ args.uploader,
67
+ )
68
+ if improved and all(upload_values):
69
+ time.sleep(5)
70
+ try:
71
+ subprocess.run(
72
+ [
73
+ str(args.uploader),
74
+ "--repo-id", args.repo_id,
75
+ "--checkpoint-prefix", str(args.checkpoint_prefix),
76
+ "--tag", f"eval-{seen:05d}",
77
+ "--metric", str(score),
78
+ "--config", str(args.config),
79
+ "--dictionary", str(args.dictionary),
80
+ ],
81
+ check=True,
82
+ )
83
+ except Exception as error:
84
+ print(f"UPLOAD_FAILED {error}", flush=True)
85
  if seen >= args.min_evals and (
86
  stale >= args.patience or diverged >= args.divergence_patience
87
  ):