ysharma HF Staff commited on
Commit
479d47a
·
verified ·
1 Parent(s): cddd103

Add training wrapper: tails ai-toolkit loss_log.db -> Trackio, uploads checkpoints to Hub

Browse files
Files changed (1) hide show
  1. trackio_tail.py +86 -0
trackio_tail.py ADDED
@@ -0,0 +1,86 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python
2
+ """Training companion process for ai-toolkit runs.
3
+
4
+ 1. Tails ai-toolkit's SQLite loss log ({save_root}/loss_log.db, schema:
5
+ metrics(step, key, value_real) with keys 'loss/...' and 'learning_rate')
6
+ and pushes each step's loss to Trackio.
7
+ 2. Watches the save_root for new {name}_{step}.safetensors LoRA checkpoints and
8
+ uploads each to the Hub repo (ai-toolkit's local save dir is rolling; the
9
+ Hub keeps every checkpoint).
10
+
11
+ Runs until a file named DONE appears next to the db, or --max-minutes expires.
12
+
13
+ Usage:
14
+ python trackio_tail.py --db /output/orbit_alpha_qwen21/loss_log.db \
15
+ --save-root /output/orbit_alpha_qwen21 \
16
+ --space-id <trackio space id> --project orbit-alpha-lora \
17
+ --hub-repo ysharma/orbit-alpha-lora [--max-minutes 300]
18
+ """
19
+ import argparse, glob, json, os, sqlite3, time
20
+
21
+ import trackio
22
+
23
+
24
+ def read_rows(db_path, last_step):
25
+ out = []
26
+ if not os.path.exists(db_path):
27
+ return out, last_step
28
+ con = sqlite3.connect(f'file:{db_path}?mode=ro', uri=True)
29
+ try:
30
+ cur = con.execute(
31
+ "SELECT step, key, value_real FROM metrics WHERE step > ? AND key LIKE 'loss/%' ORDER BY step",
32
+ (last_step,))
33
+ for step, key, val in cur:
34
+ out.append((step, key, val))
35
+ cur = con.execute("SELECT MAX(step) FROM metrics")
36
+ mx = cur.fetchone()[0] or last_step
37
+ finally:
38
+ con.close()
39
+ return out, mx
40
+
41
+
42
+ def main():
43
+ ap = argparse.ArgumentParser()
44
+ ap.add_argument('--db', required=True)
45
+ ap.add_argument('--save-root', required=True)
46
+ ap.add_argument('--space-id', required=True)
47
+ ap.add_argument('--project', required=True)
48
+ ap.add_argument('--hub-repo', required=True)
49
+ ap.add_argument('--max-minutes', type=int, default=300)
50
+ args = ap.parse_args()
51
+ trackio.init(project=args.project, space_id=args.space_id)
52
+ last_step = 0
53
+ uploaded = set()
54
+ t0 = time.time()
55
+ while time.time() - t0 < args.max_minutes * 60:
56
+ rows, mx = read_rows(args.db, last_step)
57
+ for step, key, val in rows:
58
+ trackio.log({'loss': val}, step=step)
59
+ if rows:
60
+ last_step = mx
61
+ print(f'tail: pushed loss through step {mx}', flush=True)
62
+ for ck in sorted(glob.glob(os.path.join(args.save_root, '*.safetensors'))):
63
+ if ck not in uploaded:
64
+ try:
65
+ from huggingface_hub import HfApi
66
+ HfApi(token=os.environ.get('HF_TOKEN')).upload_file(
67
+ repo_id=args.hub_repo, repo_type='model',
68
+ path_in_repo=os.path.basename(ck), path_or_fileobj=ck)
69
+ uploaded.add(ck)
70
+ print(f'tail: uploaded {os.path.basename(ck)}', flush=True)
71
+ except Exception as e:
72
+ print(f'tail: upload failed {ck}: {e}', flush=True)
73
+ uploaded.discard(ck)
74
+ if os.path.exists(os.path.join(os.path.dirname(args.db), 'DONE')):
75
+ time.sleep(20)
76
+ rows, mx = read_rows(args.db, last_step)
77
+ for step, key, val in rows:
78
+ trackio.log({'loss': val}, step=step)
79
+ break
80
+ time.sleep(10)
81
+ trackio.finish()
82
+ print('tail: done', flush=True)
83
+
84
+
85
+ if __name__ == '__main__':
86
+ main()