#!/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 --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()