ysharma's picture
ysharma HF Staff
Add training wrapper: tails ai-toolkit loss_log.db -> Trackio, uploads checkpoints to Hub
479d47a verified
Raw History Blame Contribute Delete
3.3 kB
#!/usr/bin/env python
"""Training companion process for ai-toolkit runs.
1. Tails ai-toolkit's SQLite loss log ({save_root}/loss_log.db, schema:
metrics(step, key, value_real) with keys 'loss/...' and 'learning_rate')
and pushes each step's loss to Trackio.
2. Watches the save_root for new {name}_{step}.safetensors LoRA checkpoints and
uploads each to the Hub repo (ai-toolkit's local save dir is rolling; the
Hub keeps every checkpoint).
Runs until a file named DONE appears next to the db, or --max-minutes expires.
Usage:
python trackio_tail.py --db /output/orbit_alpha_qwen21/loss_log.db \
--save-root /output/orbit_alpha_qwen21 \
--space-id <trackio space id> --project orbit-alpha-lora \
--hub-repo ysharma/orbit-alpha-lora [--max-minutes 300]
"""
import argparse, glob, json, os, sqlite3, time
import trackio
def read_rows(db_path, last_step):
out = []
if not os.path.exists(db_path):
return out, last_step
con = sqlite3.connect(f'file:{db_path}?mode=ro', uri=True)
try:
cur = con.execute(
"SELECT step, key, value_real FROM metrics WHERE step > ? AND key LIKE 'loss/%' ORDER BY step",
(last_step,))
for step, key, val in cur:
out.append((step, key, val))
cur = con.execute("SELECT MAX(step) FROM metrics")
mx = cur.fetchone()[0] or last_step
finally:
con.close()
return out, mx
def main():
ap = argparse.ArgumentParser()
ap.add_argument('--db', required=True)
ap.add_argument('--save-root', required=True)
ap.add_argument('--space-id', required=True)
ap.add_argument('--project', required=True)
ap.add_argument('--hub-repo', required=True)
ap.add_argument('--max-minutes', type=int, default=300)
args = ap.parse_args()
trackio.init(project=args.project, space_id=args.space_id)
last_step = 0
uploaded = set()
t0 = time.time()
while time.time() - t0 < args.max_minutes * 60:
rows, mx = read_rows(args.db, last_step)
for step, key, val in rows:
trackio.log({'loss': val}, step=step)
if rows:
last_step = mx
print(f'tail: pushed loss through step {mx}', flush=True)
for ck in sorted(glob.glob(os.path.join(args.save_root, '*.safetensors'))):
if ck not in uploaded:
try:
from huggingface_hub import HfApi
HfApi(token=os.environ.get('HF_TOKEN')).upload_file(
repo_id=args.hub_repo, repo_type='model',
path_in_repo=os.path.basename(ck), path_or_fileobj=ck)
uploaded.add(ck)
print(f'tail: uploaded {os.path.basename(ck)}', flush=True)
except Exception as e:
print(f'tail: upload failed {ck}: {e}', flush=True)
uploaded.discard(ck)
if os.path.exists(os.path.join(os.path.dirname(args.db), 'DONE')):
time.sleep(20)
rows, mx = read_rows(args.db, last_step)
for step, key, val in rows:
trackio.log({'loss': val}, step=step)
break
time.sleep(10)
trackio.finish()
print('tail: done', flush=True)
if __name__ == '__main__':
main()