Add training wrapper: tails ai-toolkit loss_log.db -> Trackio, uploads checkpoints to Hub
Browse files- 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()
|