AQIT 0.1.0
Loading...
Searching...
No Matches
trajectory_analysis.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2"""Multi-checkpoint training trajectory: weight-diff summary per step vs base."""
3
4from __future__ import annotations
5
6import glob
7from pathlib import Path
8from typing import Any
9
10from aquin.compute.sae_diff import load_checkpoint_state
11from aquin.compute.weight_diff import run_weight_diff
12
13
14def resolve_checkpoint_paths(*, pattern: str | None = None, directory: str | None = None) -> list[Path]:
15 paths: list[Path] = []
16 if pattern:
17 for hit in glob.glob(str(Path(pattern).expanduser()), recursive=True):
18 p = Path(hit)
19 if p.is_file():
20 paths.append(p)
21 elif directory:
22 root = Path(directory).expanduser()
23 if not root.is_dir():
24 raise FileNotFoundError(f"Directory not found: {root}")
25 paths = [p for p in root.rglob("*.pt") if p.is_file()]
26 else:
27 raise ValueError("Provide --checkpoints <glob> or --dir <path>")
28
29 if not paths:
30 raise FileNotFoundError("No checkpoint files matched.")
31
32 def _sort_key(p: Path) -> tuple[int, str]:
33 try:
34 _, step = load_checkpoint_state(p)
35 except Exception:
36 step = None
37 return (step if step is not None else 10**9, p.name)
38
39 return sorted(set(paths), key=_sort_key)
40
41
43 model_id: str,
44 checkpoint_paths: list[str | Path],
45 *,
46 name_prefix: str | None = None,
47) -> dict[str, Any]:
48 steps: list[dict[str, Any]] = []
49 prev_total: float | None = None
50
51 for path in checkpoint_paths:
52 ckpt = Path(path)
53 label = f"{name_prefix}-{ckpt.stem}" if name_prefix else ckpt.stem
54 try:
55 weight = run_weight_diff(model_id, ckpt, checkpoint_name=label)
56 except Exception as e:
57 steps.append(
58 {
59 "checkpoint": str(ckpt),
60 "name": label,
61 "error": str(e),
62 }
63 )
64 continue
65
66 if weight.get("error"):
67 steps.append(
68 {
69 "checkpoint": str(ckpt),
70 "name": label,
71 "error": weight.get("error"),
72 }
73 )
74 continue
75
76 total = float(weight.get("totalDeltaL2") or 0.0)
77 delta_from_prev = round(total - prev_total, 6) if prev_total is not None else None
78 prev_total = total
79
80 steps.append(
81 {
82 "step": weight.get("trainingStep"),
83 "checkpoint": str(ckpt),
84 "name": label,
85 "totalDeltaL2": weight.get("totalDeltaL2"),
86 "maxDeltaL2": weight.get("maxDeltaL2"),
87 "meanDeltaStableRank": weight.get("meanDeltaStableRank"),
88 "nMatrices": weight.get("nMatrices"),
89 "deltaFromPrevious": delta_from_prev,
90 "layerProfile": weight.get("layerProfile") or [],
91 "topChanged": (weight.get("topChanged") or [])[:5],
92 }
93 )
94
95 valid = [s for s in steps if not s.get("error")]
96 peak = max(valid, key=lambda s: float(s.get("totalDeltaL2") or 0), default=None)
97
98 return {
99 "schema_version": 1,
100 "type": "trajectoryAnalysis",
101 "baseModelId": model_id,
102 "nCheckpoints": len(checkpoint_paths),
103 "nAnalyzed": len(valid),
104 "steps": steps,
105 "peakStep": peak.get("step") if peak else None,
106 "peakTotalDeltaL2": peak.get("totalDeltaL2") if peak else None,
107 }
108
109
110def run_trajectory_analysis_from_args(args: dict[str, Any]) -> dict[str, Any]:
111 from aquin.compute.model_loader import get_active_model_id, resolve_model_id
112
113 model_id = args.get("model_id") or get_active_model_id() or "llama-3.2-1b"
114 try:
115 model_id = resolve_model_id(str(model_id))
116 except ValueError as e:
117 return {"error": str(e)}
118
119 try:
121 pattern=args.get("checkpoints"),
122 directory=args.get("dir"),
123 )
124 except (FileNotFoundError, ValueError) as e:
125 return {"error": str(e)}
126
128 model_id,
129 [str(p) for p in paths],
130 name_prefix=args.get("name"),
131 )
list[Path] resolve_checkpoint_paths(*, str|None pattern=None, str|None directory=None)
dict[str, Any] run_trajectory_analysis_from_args(dict[str, Any] args)
dict[str, Any] run_trajectory_analysis(str model_id, list[str|Path] checkpoint_paths, *, str|None name_prefix=None)