AQIT 0.1.0
Loading...
Searching...
No Matches
weight_trojans.py
Go to the documentation of this file.
1# Copyright (c) 2025-present Aquin Labs Private Limited. All Rights Reserved.
2# This file is part of the Aquin Engine. Unauthorized copying, modification,
3# distribution, or use of this file, via any medium, is strictly prohibited.
4# Proprietary and confidential. See LICENSE for terms.
5
6"""
7Ingested from inspection-backend/weight_trojans.py.
8Import adaptation: causal_trace -> aquin.compute.model_loader.
9FastAPI router stripped — run_weight_trojan_analysis() is the direct entry point.
10"""
11from __future__ import annotations
12
13from typing import Optional
14
15import math
16import numpy as np
17import torch
18
19KURTOSIS_FLAG = 4.0
20OUTLIER_FLAG = 0.002
21SV_RATIO_FLAG = 8.0
22RISK_SUSPICIOUS = 0.35
23RISK_HIGH = 0.65
25SIGNAL_WEIGHTS = {"kurtosis": 0.40, "outlier": 0.35, "sv_ratio": 0.25}
28def _excess_kurtosis(t: torch.Tensor) -> float:
29 f = t.float().flatten()
30 n = f.numel()
31 if n < 8:
32 return 0.0
33 mu = f.mean()
34 sigma = f.std()
35 if sigma < 1e-9:
36 return 0.0
37 z = (f - mu) / sigma
38 kurt = float((z ** 4).mean()) - 3.0
39 return kurt
40
41
42def _outlier_density(t: torch.Tensor) -> float:
43 f = t.float().flatten()
44 mu = f.mean()
45 sigma = f.std()
46 if sigma < 1e-9:
47 return 0.0
48 threshold = mu + 4.0 * sigma
49 return float((f.abs() > threshold.abs()).float().mean())
50
51
52def _sv_ratio(t: torch.Tensor) -> float:
53 f = t.float()
54 if f.dim() == 1:
55 return 1.0
56 if f.dim() > 2:
57 f = f.reshape(-1, f.shape[-1])
58 if min(f.shape) < 2:
59 return 1.0
60 # Power iteration on CPU — avoids extra GPU allocations while model is loaded.
61 f = f.cpu()
62 try:
63 v = torch.randn(f.shape[1])
64 for _ in range(20):
65 u = f @ v
66 u = u / u.norm().clamp(min=1e-9)
67 v = f.T @ u
68 v = v / v.norm().clamp(min=1e-9)
69 sigma1 = (u.unsqueeze(0) @ f @ v.unsqueeze(1)).reshape(-1)[0].item()
70
71 f2 = f - sigma1 * u.unsqueeze(1) * v.unsqueeze(0)
72 v2 = torch.randn(f.shape[1])
73 for _ in range(20):
74 u2 = f2 @ v2
75 u2 = u2 / u2.norm().clamp(min=1e-9)
76 v2 = f2.T @ u2
77 v2 = v2 / v2.norm().clamp(min=1e-9)
78 sigma2 = (u2.unsqueeze(0) @ f2 @ v2.unsqueeze(1)).reshape(-1)[0].item()
79
80 if abs(sigma2) < 1e-6:
81 return abs(sigma1) / 1e-6
82 return abs(sigma1) / max(abs(sigma2), 1e-6)
83 except Exception:
84 return 1.0
85
86
87def _analyse_tensor(name: str, t: torch.Tensor) -> dict:
88 kurt = _excess_kurtosis(t)
89 out = _outlier_density(t)
90 svr = _sv_ratio(t)
91 return {
92 "name": name,
93 "shape": list(t.shape),
94 "n_params": t.numel(),
95 "mean": round(float(t.float().mean()), 6),
96 "std": round(float(t.float().std()), 6),
97 "kurtosis": round(kurt, 4),
98 "outlier_density": round(out, 6),
99 "sv_ratio": round(svr, 4),
100 }
101
102
103def _compute_risk(layer_stats: list[dict]) -> list[dict]:
104 if not layer_stats:
105 return []
106
107 kurtosis_vals = [s["kurtosis"] for s in layer_stats]
108 outlier_vals = [s["outlier_density"] for s in layer_stats]
109 sv_vals = [s["sv_ratio"] for s in layer_stats]
110
111 def _zscore_norm(vals: list[float]) -> list[float]:
112 arr = np.array(vals, dtype=float)
113 mu, sigma = arr.mean(), arr.std()
114 if sigma < 1e-9:
115 return [0.0] * len(vals)
116 return np.clip((arr - mu) / sigma, 0, None).tolist()
117
118 kurt_z = _zscore_norm(kurtosis_vals)
119 out_z = _zscore_norm(outlier_vals)
120 sv_z = _zscore_norm(sv_vals)
121
122 def _norm01(vals: list[float]) -> list[float]:
123 mx = max(vals) if vals else 0.0
124 if mx < 1e-9:
125 return [0.0] * len(vals)
126 return [v / mx for v in vals]
127
128 kurt_n = _norm01(kurt_z)
129 out_n = _norm01(out_z)
130 sv_n = _norm01(sv_z)
131
132 wk = SIGNAL_WEIGHTS["kurtosis"]
133 wo = SIGNAL_WEIGHTS["outlier"]
134 ws = SIGNAL_WEIGHTS["sv_ratio"]
135
136 results = []
137 for i, s in enumerate(layer_stats):
138 risk = wk * kurt_n[i] + wo * out_n[i] + ws * sv_n[i]
139 risk = round(min(risk, 1.0), 4)
140 status = (
141 "high_risk" if risk >= RISK_HIGH else
142 "suspicious" if risk >= RISK_SUSPICIOUS else
143 "clean"
144 )
145 flags = []
146 if s["kurtosis"] > KURTOSIS_FLAG:
147 flags.append(f"excess kurtosis {s['kurtosis']:.2f} (heavy-tailed weight distribution)")
148 if s["outlier_density"] > OUTLIER_FLAG:
149 flags.append(f"outlier density {s['outlier_density']:.4%} (weight magnitude spikes)")
150 if s["sv_ratio"] > SV_RATIO_FLAG:
151 flags.append(f"SV ratio {s['sv_ratio']:.1f} (low-rank implant signature)")
152
153 results.append({**s, "risk_score": risk, "status": status, "flags": flags})
154
155 return results
156
157
159 model,
160 layer_range: Optional[list[int]] = None,
161 max_tensors_per_layer: int = 4,
162) -> dict:
163 from aquin.compute.weight_rank import _is_heavy_weight_model, _should_skip_weight_param
164
165 heavy = _is_heavy_weight_model(model)
166 if heavy:
167 max_tensors_per_layer = min(max_tensors_per_layer, 6)
168
169 target_layers = set(layer_range) if layer_range else None
170 layer_stats: list[dict] = []
171 skipped = 0
172 per_layer_count: dict[int, int] = {}
173
174 for name, param in model.named_parameters():
175 parts = name.split(".")
176 layer_idx: Optional[int] = None
177 for p in parts:
178 if p.isdigit():
179 layer_idx = int(p)
180 break
181
182 if target_layers is not None and layer_idx not in target_layers:
183 skipped += 1
184 continue
185
186 if param.dim() < 2 or param.numel() < 256:
187 continue
188
189 if _should_skip_weight_param(name, heavy=heavy):
190 skipped += 1
191 continue
192
193 if layer_idx is not None:
194 count = per_layer_count.get(layer_idx, 0)
195 if count >= max_tensors_per_layer:
196 skipped += 1
197 continue
198 per_layer_count[layer_idx] = count + 1
199
200 stats = _analyse_tensor(name, param.detach().cpu())
201 if layer_idx is not None:
202 stats["layer_idx"] = layer_idx
203 layer_stats.append(stats)
204
205 if not layer_stats:
206 return {"error": "No eligible weight tensors found", "layers_analysed": 0}
207
208 scored = _compute_risk(layer_stats)
209
210 high_risk = [s for s in scored if s["status"] == "high_risk"]
211 suspicious = [s for s in scored if s["status"] == "suspicious"]
212 clean = [s for s in scored if s["status"] == "clean"]
213
214 n = len(scored)
215 pct_flagged = round((len(high_risk) + len(suspicious)) / max(n, 1) * 100, 2)
216
217 if len(high_risk) >= 2:
218 verdict = "high_risk"
219 elif len(high_risk) >= 1 or len(suspicious) >= 3:
220 verdict = "suspicious"
221 else:
222 verdict = "clean"
223
224 composite = round(float(np.mean([s["risk_score"] for s in scored])), 4) if scored else 0.0
225
226 all_flags: list[str] = []
227 for s in high_risk + suspicious:
228 for f in s["flags"]:
229 if f not in all_flags:
230 all_flags.append(f)
231
232 return {
233 "layers_analysed": n,
234 "tensors_skipped": skipped,
235 "composite_risk": composite,
236 "verdict": verdict,
237 "pct_flagged": pct_flagged,
238 "high_risk_count": len(high_risk),
239 "suspicious_count": len(suspicious),
240 "clean_count": len(clean),
241 "all_flags": all_flags,
242 "scored_tensors": scored,
243 "signals": list(SIGNAL_WEIGHTS.keys()),
244 "heavy_model": heavy,
245 "thresholds": {
246 "kurtosis_flag": KURTOSIS_FLAG,
247 "outlier_flag": OUTLIER_FLAG,
248 "sv_ratio_flag": SV_RATIO_FLAG,
249 "risk_suspicious": RISK_SUSPICIOUS,
250 "risk_high": RISK_HIGH,
251 },
252 }
list[dict] _compute_risk(list[dict] layer_stats)
float _sv_ratio(torch.Tensor t)
dict _analyse_tensor(str name, torch.Tensor t)
dict run_weight_trojan_analysis(model, Optional[list[int]] layer_range=None, int max_tensors_per_layer=4)
float _outlier_density(torch.Tensor t)
float _excess_kurtosis(torch.Tensor t)