Download trackio_tail.py from ML-Intern-lab/Qwen-Image-2.1-viewpoint-orbit-LoRA: direct link, hf CLI and curl.
- Browser
- Download file 3.3 kB
-
https://huggingface.co/ML-Intern-lab/Qwen-Image-2.1-viewpoint-orbit-LoRA/resolve/main/trackio_tail.py
- Command line
-
hf download hf://ML-Intern-lab/Qwen-Image-2.1-viewpoint-orbit-LoRA/trackio_tail.py
-
curl -L -o trackio_tail.py https://huggingface.co/ML-Intern-lab/Qwen-Image-2.1-viewpoint-orbit-LoRA/resolve/main/trackio_tail.py
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() |