Spaces:
Sleeping
Sleeping
File size: 4,435 Bytes
17ddbea bfd323f 17ddbea bfd323f 17ddbea bfd323f 17ddbea bfd323f 17ddbea | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 | import { describe, expect, it } from 'vitest';
import {
DEFAULT_SAE_SETTINGS,
computeActivationMetrics,
formatSaeTrainProgress,
loadSaeSettings,
resolveSaeSettings,
saveSaeSettings,
} from '../src/ui/saeControlsDefaults.js';
import {
applySaeToCompare,
cosineSimilarity,
densifyTopKActivations,
} from '../src/core/saeReplace.js';
describe('saeControlsDefaults', () => {
it('round-trips enabled + train params through localStorage', () => {
const store = new Map();
const storage = {
getItem: (k) => (store.has(k) ? store.get(k) : null),
setItem: (k, v) => { store.set(k, String(v)); },
removeItem: (k) => { store.delete(k); },
key: (i) => [...store.keys()][i] ?? null,
get length() { return store.size; },
};
const settings = resolveSaeSettings({
enabled: true,
hiddenDim: 4096,
k: 16,
epochs: 10,
lr: 0.002,
batchSize: 32,
});
saveSaeSettings(settings, storage);
const loaded = loadSaeSettings(storage);
expect(loaded.enabled).toBe(true);
expect(loaded.hiddenDim).toBe(4096);
expect(loaded.k).toBe(16);
expect(loaded.epochs).toBe(10);
expect(loaded.batchSize).toBe(32);
});
it('defaults are 8192 / k32 / 20ep and purge legacy poisoned keys', () => {
expect(DEFAULT_SAE_SETTINGS.enabled).toBe(false);
expect(DEFAULT_SAE_SETTINGS.hiddenDim).toBe(8192);
expect(DEFAULT_SAE_SETTINGS.k).toBe(32);
expect(DEFAULT_SAE_SETTINGS.epochs).toBe(20);
expect(loadSaeSettings(null).enabled).toBe(false);
const store = new Map([
['vl3d.sae.hiddenDim', '32'],
['vl3d.sae.k', '1'],
['vl3d.sae.epochs', '1'],
['vl3d.sae.enabled', 'true'],
]);
const storage = {
getItem: (k) => (store.has(k) ? store.get(k) : null),
setItem: (k, v) => { store.set(k, String(v)); },
removeItem: (k) => { store.delete(k); },
key: (i) => [...store.keys()][i] ?? null,
get length() { return store.size; },
};
const loaded = loadSaeSettings(storage);
expect(loaded.hiddenDim).toBe(8192);
expect(loaded.k).toBe(32);
expect(loaded.epochs).toBe(20);
expect(loaded.enabled).toBe(false);
expect(store.has('vl3d.sae.hiddenDim')).toBe(false);
});
it('computes L0 / sparsity from sparse rows', () => {
const acts = [
[1, 0, 0, 2],
[0, 0, 3, 0],
];
const m = computeActivationMetrics(acts);
expect(m.dim).toBe(4);
expect(m.l0).toBe(1.5);
expect(m.activeFeatures).toBe(3);
});
it('formats train progress with done / left / percent', () => {
const mid = formatSaeTrainProgress({
status: 'training',
phase_key: 'training',
message: 'Training epoch 13/50 — 38 remaining · last loss=0.012345',
current_epoch: 12,
total_epochs: 50,
remaining_epochs: 38,
percent: 24,
n_vectors: 40,
resolved_hidden: 160,
resolved_k: 16,
});
expect(mid.busy).toBe(true);
expect(mid.label).toContain('38 remaining');
expect(mid.meta).toContain('12/50 done');
expect(mid.meta).toContain('38 left');
expect(mid.meta).toContain('24%');
expect(mid.percent).toBe(24);
const prep = formatSaeTrainProgress({
status: 'training',
phase_key: 'preparing',
message: '',
current_epoch: 0,
total_epochs: 50,
});
expect(prep.label).toMatch(/Preparing/i);
expect(prep.meta).toContain('0/50 done');
expect(prep.meta).toContain('50 left');
expect(prep.indeterminate).toBe(true);
});
});
describe('saeReplace', () => {
it('recomputes compare cosine_vs_first in SAE space', () => {
const raw = {
count: 2,
items: [
{ id: 'tok_0', text: 'a', embedding: [1, 0] },
{ id: 'tok_1', text: 'b', embedding: [0, 1] },
],
};
const acts = [
[1, 0, 0],
[1, 0, 0],
];
const next = applySaeToCompare(raw, acts);
expect(next.items[0].cosine_vs_first).toBe(1);
expect(next.items[1].cosine_vs_first).toBeCloseTo(1, 5);
expect(cosineSimilarity([1, 0], [0, 1])).toBeCloseTo(0, 5);
});
it('densifies Top-K sparse encode payload', () => {
const dense = densifyTopKActivations({
format: 'topk_sparse',
indices: [[0, 2], [1, 3]],
values: [[1.5, 0.5], [2, 0]],
dimension: 4,
k: 2,
});
expect(dense).toEqual([
[1.5, 0, 0.5, 0],
[0, 2, 0, 0],
]);
});
});
|