AQIT 0.1.0
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
4from __future__ import annotations
5
6import sys
7from typing import Any
8
9
10def _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
17def _has_flag(args: list[str], name: str) -> bool:
18 return name in args
19
20
21def 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
36def _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
47def cmd_sweep(args: list[str]) -> None:
48 if not args or args[0] in ("-h", "--help", "help"):
50 return
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}")
list[float] parse_strengths(str|None raw)
Definition sweep_cli.py:25
bool _has_flag(list[str] args, str name)
Definition sweep_cli.py:21
None _print_help()
Definition sweep_cli.py:40
None cmd_sweep(list[str] args)
Definition sweep_cli.py:51
str|None _parse_flag(list[str] args, str name)
Definition sweep_cli.py:14