Download spectra_rsi/probes.py from kiruluta/SPECTRA-RSI-HF-Scaling-Benchmark: direct link, hf CLI and curl.
- Browser
- Download file 2.64 kB
-
https://huggingface.co/kiruluta/SPECTRA-RSI-HF-Scaling-Benchmark/resolve/main/spectra_rsi/probes.py
- Command line
-
hf download hf://kiruluta/SPECTRA-RSI-HF-Scaling-Benchmark/spectra_rsi/probes.py
-
curl -L -o probes.py https://huggingface.co/kiruluta/SPECTRA-RSI-HF-Scaling-Benchmark/resolve/main/spectra_rsi/probes.py
2.64 kB
| """Counterfactual spectral response dictionary (Sec. 4). | |
| Signed central-difference probes on expert directions estimate the | |
| interventional response matrix H_t; per-expert truncated SVD yields the | |
| grouped dictionary Psi_t = [U_1, ..., U_E]. | |
| """ | |
| from dataclasses import dataclass | |
| import numpy as np | |
| class Dictionary: | |
| Psi: np.ndarray # (n, p) column-orthonormal within groups | |
| groups: list # list of index arrays into columns of Psi | |
| H: np.ndarray # (n, d) raw fingerprint matrix | |
| expert_maps: list # per-expert (Sigma V^T) mapping dict coords -> beta coords | |
| def p(self): | |
| return self.Psi.shape[1] | |
| def build_dictionary(world, cfg, rng=None) -> Dictionary: | |
| """Apply +/- h probes for every expert direction and factorize per expert.""" | |
| rng = rng or np.random.default_rng(cfg.seed) | |
| E, R = cfg.n_experts, cfg.rank_per_expert | |
| d = E * R | |
| cols = [] | |
| for idx in range(d): | |
| direction = np.zeros(d) | |
| direction[idx] = 1.0 | |
| cols.append(world.probe_delta(direction, cfg.probe_magnitude, | |
| cfg.probe_items, rng=rng)) | |
| H = np.stack(cols, axis=1) # (n, d) | |
| Psi_blocks, groups, expert_maps = [], [], [] | |
| col0 = 0 | |
| for e in range(E): | |
| He = H[:, e * R:(e + 1) * R] # (n, R) | |
| U, S, Vt = np.linalg.svd(He, full_matrices=False) | |
| r = min(R, (S > 1e-8).sum()) | |
| r = max(r, 1) | |
| Psi_blocks.append(U[:, :r]) | |
| groups.append(np.arange(col0, col0 + r)) | |
| expert_maps.append((S[:r], Vt[:r])) | |
| col0 += r | |
| Psi = np.concatenate(Psi_blocks, axis=1) | |
| return Dictionary(Psi=Psi, groups=groups, H=H, expert_maps=expert_maps) | |
| def pilot_linearity_check(world, cand, dictionary, cfg, rng=None): | |
| """Trust-region pilot (Assumption 1): compare probe-superposition | |
| prediction H beta with a small paired evaluation on random slices. | |
| Returns (relative_residual, passed).""" | |
| rng = rng or np.random.default_rng(cfg.seed + 7) | |
| pred = dictionary.H @ cand.beta | |
| # test where the superposition prediction claims signal exists: the | |
| # top-|pred| slices; a linearity violation shows up exactly there. | |
| k = min(world.n, 40) | |
| idx = np.argsort(np.abs(pred))[::-1][:k] | |
| obs = np.array([world.paired_scores(cand, int(j), cfg.pilot_items, | |
| rng=rng).mean() for j in idx]) | |
| denom = max(np.linalg.norm(obs), np.linalg.norm(pred[idx])) + 1e-9 | |
| rel = np.linalg.norm(obs - pred[idx]) / denom | |
| return rel, rel <= cfg.trust_region_residual | |