AQIT
0.1.0
Toggle main menu visibility
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
4
from
__future__
import
annotations
5
6
import
glob
7
from
pathlib
import
Path
8
from
typing
import
Any
9
10
from
aquin.compute.sae_diff
import
load_checkpoint_state
11
from
aquin.compute.weight_diff
import
run_weight_diff
12
13
14
def
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
42
def
run_trajectory_analysis
(
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
110
def
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
:
120
paths =
resolve_checkpoint_paths
(
121
pattern=args.get(
"checkpoints"
),
122
directory=args.get(
"dir"
),
123
)
124
except
(FileNotFoundError, ValueError)
as
e:
125
return
{
"error"
: str(e)}
126
127
return
run_trajectory_analysis
(
128
model_id,
129
[str(p)
for
p
in
paths],
130
name_prefix=args.get(
"name"
),
131
)
aquin.compute.model_loader
Definition
model_loader.py:1
aquin.compute.sae_diff
Definition
sae_diff.py:1
aquin.compute.trajectory_analysis.resolve_checkpoint_paths
list[Path] resolve_checkpoint_paths(*, str|None pattern=None, str|None directory=None)
Definition
trajectory_analysis.py:18
aquin.compute.trajectory_analysis.run_trajectory_analysis_from_args
dict[str, Any] run_trajectory_analysis_from_args(dict[str, Any] args)
Definition
trajectory_analysis.py:114
aquin.compute.trajectory_analysis.run_trajectory_analysis
dict[str, Any] run_trajectory_analysis(str model_id, list[str|Path] checkpoint_paths, *, str|None name_prefix=None)
Definition
trajectory_analysis.py:51
aquin.compute.weight_diff
Definition
weight_diff.py:1
aquin
compute
trajectory_analysis.py
AQIT · Aquin Labs Private Limited · Apache 2.0 · Generated by
1.18.0