AQIT
0.1.0
Toggle main menu visibility
Loading...
Searching...
No Matches
sweep_cli.py
Go to the documentation of this file.
1
# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2
"""aquin sweep — steering sweep across strengths."""
3
4
from
__future__
import
annotations
5
6
import
sys
7
from
typing
import
Any
8
9
10
def
_parse_flag
(args: list[str], name: str) -> str |
None
:
11
for
i, a
in
enumerate(args):
12
if
a == name
and
i + 1 < len(args):
13
return
args[i + 1]
14
return
None
15
16
17
def
_has_flag
(args: list[str], name: str) -> bool:
18
return
name
in
args
19
20
21
def
parse_strengths
(raw: str |
None
) -> list[float]:
22
text = (raw
or
""
).strip()
23
if
not
text:
24
return
[-10.0, -5.0, 0.0, 5.0, 10.0]
25
out: list[float] = []
26
for
part
in
text.split(
","
):
27
part = part.strip()
28
if
not
part:
29
continue
30
out.append(float(part))
31
if
not
out:
32
raise
ValueError(
"No strengths parsed"
)
33
return
out
34
35
36
def
_print_help
() -> None:
37
print(
"Sweep steering strength over a fixed feature/vector."
)
38
print(
""
)
39
print(
"Usage: aquin sweep (--feature_idx <n> | --vector <path>) [--strengths <csv>]"
)
40
print(
" [--prompt <text>] [--eval] [--prompts <json|jsonl>] [--reference_answers <json>]"
)
41
print(
""
)
42
print(
"Examples:"
)
43
print(
' aquin sweep --feature_idx 42 --prompt "Explain photosynthesis"'
)
44
print(
" aquin sweep --vector vec.json --eval --prompts probes.jsonl"
)
45
46
47
def
cmd_sweep
(args: list[str]) ->
None
:
48
if
not
args
or
args[0]
in
(
"-h"
,
"--help"
,
"help"
):
49
_print_help
()
50
return
51
52
from
aquin.compute.loader_shim
import
apply
as
_shim_apply
53
from
aquin.engine.local_server
import
start
as
_start_local_server
54
from
aquin.compute.model_loader
import
(
55
get_active_model_id,
56
resolve_model_id,
57
)
58
from
aquin.compute.steer_vector
import
run_steer_with_vector
59
60
feature_idx_raw =
_parse_flag
(args,
"--feature_idx"
)
61
vector =
_parse_flag
(args,
"--vector"
)
62
if
feature_idx_raw
is
None
and
not
vector:
63
print(
"Error: pass --feature_idx <n> or --vector <path>"
)
64
sys.exit(1)
65
try
:
66
strengths =
parse_strengths
(
_parse_flag
(args,
"--strengths"
))
67
except
ValueError
as
e:
68
print(f
"Error: {e}"
)
69
sys.exit(1)
70
71
active = (get_active_model_id()
or
""
).strip()
72
if
not
active:
73
print(
"Error: no model loaded. Run: aquin load model <id>"
)
74
sys.exit(1)
75
_shim_apply()
76
_start_local_server()
77
78
model_id = resolve_model_id(active)
79
feature_idx = int(feature_idx_raw)
if
feature_idx_raw
is
not
None
else
None
80
prompt =
_parse_flag
(args,
"--prompt"
)
81
layer_raw =
_parse_flag
(args,
"--layer"
)
82
layer = int(layer_raw)
if
layer_raw
else
None
83
max_new_tokens_raw =
_parse_flag
(args,
"--max_new_tokens"
)
84
max_new_tokens = int(max_new_tokens_raw)
if
max_new_tokens_raw
else
80
85
do_eval =
_has_flag
(args,
"--eval"
)
86
prompts =
_parse_flag
(args,
"--prompts"
)
87
refs =
_parse_flag
(args,
"--reference_answers"
)
88
threshold =
_parse_flag
(args,
"--threshold"
)
89
max_probes =
_parse_flag
(args,
"--max_probes"
)
90
91
rows: list[dict[str, Any]] = []
92
feature_label =
None
93
resolved_layer = layer
94
for
strength
in
strengths:
95
tool_args: dict[str, Any] = {
96
"prompt"
: prompt,
97
"eval"
: do_eval,
98
"prompts"
: prompts,
99
"reference_answers"
: refs,
100
}
101
if
threshold
is
not
None
:
102
tool_args[
"threshold"
] = float(threshold)
103
if
max_probes
is
not
None
:
104
tool_args[
"max_probes"
] = int(max_probes)
105
result = run_steer_with_vector(
106
model_id=model_id,
107
prompt=prompt,
108
steer_strength=float(strength),
109
layer=resolved_layer,
110
feature_idx=feature_idx,
111
vector_path=vector,
112
feature_label=feature_label,
113
max_new_tokens=max_new_tokens,
114
args=tool_args,
115
)
116
if
result.get(
"error"
):
117
print(f
"Error at strength {strength}: {result['error']}"
)
118
sys.exit(1)
119
feature_label = str(result.get(
"feature_label"
)
or
feature_label
or
""
)
120
resolved_layer = int(result.get(
"layer"
)
or
resolved_layer
or
0)
121
row: dict[str, Any] = {
122
"strength"
: float(strength),
123
"prompt"
: result.get(
"prompt"
),
124
"steered_response"
: result.get(
"steered_response"
),
125
}
126
eval_data = result.get(
"eval"
)
if
isinstance(result.get(
"eval"
), dict)
else
None
127
if
eval_data:
128
row[
"pass_rate"
] = (eval_data.get(
"steered"
)
or
{}).get(
"pass_rate"
)
129
row[
"delta_pass_rate"
] = eval_data.get(
"delta_pass_rate"
)
130
row[
"mode"
] = eval_data.get(
"mode"
)
131
rows.append(row)
132
133
print(
134
f
"[sweep] model={model_id} feature={feature_idx if feature_idx is not None else vector} "
135
f
"layer={resolved_layer if resolved_layer is not None else '—'} n={len(rows)}"
136
)
137
if
do_eval:
138
print(
""
)
139
print(
"strength pass delta mode"
)
140
for
row
in
rows:
141
pr = row.get(
"pass_rate"
)
142
dpr = row.get(
"delta_pass_rate"
)
143
pr_s =
"—"
if
pr
is
None
else
f
"{100.0 * float(pr):5.1f}%"
144
dpr_s =
"—"
if
dpr
is
None
else
f
"{100.0 * float(dpr):+5.1f}%"
145
print(f
"{row['strength']:>7.2f} {pr_s:>6} {dpr_s:>7} {row.get('mode', '—')}"
)
146
else
:
147
print(
""
)
148
print(
"strength response"
)
149
for
row
in
rows:
150
resp =
" "
.join(str(row.get(
"steered_response"
)
or
""
).split())
151
if
len(resp) > 88:
152
resp = resp[:87] +
"…"
153
print(f
"{row['strength']:>7.2f} {resp}"
)
aquin.compute.loader_shim
Definition
loader_shim.py:1
aquin.compute.model_loader
Definition
model_loader.py:1
aquin.compute.steer_vector
Definition
steer_vector.py:1
aquin.engine.local_server
Definition
local_server.py:1
aquin.sweep_cli.parse_strengths
list[float] parse_strengths(str|None raw)
Definition
sweep_cli.py:25
aquin.sweep_cli._has_flag
bool _has_flag(list[str] args, str name)
Definition
sweep_cli.py:21
aquin.sweep_cli._print_help
None _print_help()
Definition
sweep_cli.py:40
aquin.sweep_cli.cmd_sweep
None cmd_sweep(list[str] args)
Definition
sweep_cli.py:51
aquin.sweep_cli._parse_flag
str|None _parse_flag(list[str] args, str name)
Definition
sweep_cli.py:14
aquin
sweep_cli.py
AQIT · Aquin Labs Private Limited · Apache 2.0 · Generated by
1.18.0