AQIT 0.1.0
Loading...
Searching...
No Matches
trajectory_analysis_cli.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2"""aquin check trajectory — multi-checkpoint weight trajectory vs base."""
3
4import sys
5
6from aquin.cli_flags import reject_legacy_output_flags
7from pathlib import Path
8from typing import Any
9
10
11def _parse_flag(args: list[str], name: str) -> str | None:
12 for i, a in enumerate(args):
13 if a == name and i + 1 < len(args):
14 return args[i + 1]
15 return None
16
17
18def _has_flag(args: list[str], name: str) -> bool:
19 return name in args
20
21
22def _ensure_compute_env() -> None:
23 from aquin.compute.loader_shim import apply as _shim_apply
24 from aquin.engine.local_server import start as _start_local_server
25
26 _shim_apply()
27 _start_local_server()
28
29
30def _resolve_model_id() -> str:
31 from aquin.compute.model_loader import get_active_model_id, resolve_model_id
32
33 active = (get_active_model_id() or "").strip()
34 if not active:
35 print("Error: no model loaded. Run: aquin load --model <id>")
36 sys.exit(1)
37 try:
38 return resolve_model_id(active)
39 except ValueError as e:
40 print(f"Error: {e}")
41 sys.exit(1)
42
43
44def _print_help() -> None:
45 print("Training trajectory: weight-diff summary for each checkpoint vs base.")
46 print("")
47 print("Prerequisite: aquin load --model <id>")
48 print("")
49 print("Usage: aquin check trajectory --checkpoints <glob>")
50 print(" aquin check trajectory --dir <checkpoint-dir>")
51 print(" [--name <prefix>] [--save <path>]")
52 print("")
53 print(" --checkpoints Glob of .pt checkpoints (sorted by training step when present).")
54 print(" --dir Scan directory recursively for *.pt files.")
55 print(" --name Optional prefix for step labels.")
56 print(" --save Write schema_version=1 JSON export.")
57 print("")
58 print("Example:")
59 print(" aquin check trajectory --dir ~/runs/checkpoints --name my-run")
60 print("")
61 print("Docs: https://aquin.app/docs/checkpoint-sae")
62
63
64def cmd_trajectory_analysis(args: list[str]) -> None:
65 if _has_flag(args, "--help") or _has_flag(args, "-h"):
67 return
69 reject_legacy_output_flags(args)
70
71 pattern = _parse_flag(args, "--checkpoints")
72 directory = _parse_flag(args, "--dir")
73 if not pattern and not directory:
74 print("Error: --checkpoints <glob> or --dir <path> is required.")
75 print("")
77 sys.exit(1)
78
79 if _parse_flag(args, "--model") is not None:
80 print("Error: check trajectory uses the loaded session model only.")
81 sys.exit(1)
82
84
85 from aquin.cli import _build_tool_ctx
86 from aquin.engine.sync_dispatch import dispatch_with_sync, require_active_session
87
88 mid = _resolve_model_id()
89 ctx = _build_tool_ctx(model_id=mid)
90 require_active_session(ctx, label="aquin check trajectory")
91
92 tool_args: dict[str, Any] = {
93 "model_id": mid,
94 "checkpoints": pattern,
95 "dir": directory,
96 "name": _parse_flag(args, "--name"),
97 "save": _parse_flag(args, "--save"),
98 }
99 try:
100 label = pattern or directory or ""
101 print(f"[check trajectory] model={mid} source={label}")
102 result = dispatch_with_sync("run_trajectory_analysis", tool_args, ctx)
103 except Exception as e:
104 print(f"Error: {e}")
105 sys.exit(1)
106
107 from aquin.cli_output import print_tool_result
108
109 print_tool_result("check trajectory", result)
110
111 if isinstance(result, dict) and result.get("error"):
112 sys.exit(1)
bool _has_flag(list[str] args, str name)
str|None _parse_flag(list[str] args, str name)
None cmd_trajectory_analysis(list[str] args)