cofiber-detection / circuit /evolve_fast.py
phanerozoic's picture
update repository
dbbceb8
Raw
History Blame Contribute Delete
7.33 kB
"""
Fixed-K batched GPU evolution. All individuals have exactly K dims.
Genome: (POP, K) int tensor — indices into 768 feature dims.
Fitness: one batched torch.linalg.solve over (POP, K+1, K+1).
Target: hundreds of gen/s.
"""
import json, os, sys, time
import torch
SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
REPO_ROOT = os.path.dirname(SCRIPT_DIR)
VAL_TENSORS = os.path.join(REPO_ROOT, "analytical_stats_cache", "val_tensors_layernorm_500.pt")
GREEDY_PATH = os.path.join(REPO_ROOT, "analytical_variants", "greedy_forward_gpu.json")
PERSON_CLASS = 0
DEVICE = "cuda"
@torch.no_grad()
def batch_fitness(features, is_person, pop_dims, lam=0.1):
"""Evaluate entire population in one batched solve.
features: (N, 768)
is_person: (N,) bool
pop_dims: (POP, K) long — which dims each individual uses
Returns: (POP,) F1 scores
"""
POP, K = pop_dims.shape
N = features.shape[0]
# Gather features for all individuals: (POP, N, K)
f_batch = features[:, :].unsqueeze(0).expand(POP, -1, -1) # (POP, N, 768)
idx = pop_dims.unsqueeze(1).expand(-1, N, -1) # (POP, N, K)
f_sub = torch.gather(f_batch, 2, idx) # (POP, N, K)
# Augment with bias column: (POP, N, K+1)
ones = torch.ones(POP, N, 1, device=DEVICE)
fa = torch.cat([f_sub, ones], dim=2)
# XtX: (POP, K+1, K+1) = fa^T @ fa
XtX = torch.bmm(fa.transpose(1, 2), fa)
# Regularize
I = torch.eye(K + 1, device=DEVICE).unsqueeze(0).expand(POP, -1, -1)
XtX = XtX + lam * I * N
# XtY: (POP, K+1, 1) = fa^T @ y
y = is_person.float().unsqueeze(0).unsqueeze(2).expand(POP, -1, -1) # (POP, N, 1)
XtY = torch.bmm(fa.transpose(1, 2), y)
# Batched solve: (POP, K+1, 1)
try:
W = torch.linalg.solve(XtX, XtY)
except Exception:
return torch.zeros(POP, device=DEVICE)
# Predict: (POP, N, 1)
scores = torch.bmm(fa, W)
pred = scores.squeeze(2) > 0.5 # (POP, N)
# F1 per individual
is_p = is_person.unsqueeze(0).expand(POP, -1)
tp = (pred & is_p).sum(dim=1).float()
fp = (pred & ~is_p).sum(dim=1).float()
fn = (~pred & is_p).sum(dim=1).float()
prec = tp / (tp + fp).clamp(min=1)
rec = tp / (tp + fn).clamp(min=1)
f1 = 2 * prec * rec / (prec + rec).clamp(min=1e-9)
return f1
def main():
print("=" * 60)
print("Fixed-K Batched GPU Evolution")
print("=" * 60, flush=True)
val = torch.load(VAL_TENSORS, map_location="cpu", weights_only=False)
features = val["features"]
is_person = (val["cls_targets"] == PERSON_CLASS)
pos_idx = is_person.nonzero(as_tuple=True)[0]
neg_idx = (~is_person).nonzero(as_tuple=True)[0]
n_take = min(2000, len(pos_idx))
sel = torch.cat([pos_idx[torch.randperm(len(pos_idx))[:n_take]],
neg_idx[torch.randperm(len(neg_idx))[:n_take]]])
sel = sel[torch.randperm(len(sel))]
sub_f = features[sel].to(DEVICE)
sub_person = is_person[sel].to(DEVICE)
N = len(sel)
print(f" {N} vectors on {DEVICE}", flush=True)
greedy_dims = list(range(100))
if os.path.isfile(GREEDY_PATH):
with open(GREEDY_PATH) as f:
greedy_dims = json.load(f)["selected_dims"]
POP = 512
GEN = 5000
ELITE = 30
TARGETS = [10, 20, 50, 100, 200]
all_results = []
for K in TARGETS:
print(f"\n{'='*60}")
print(f" K={K} dims | pop={POP} | gen={GEN}")
print(f"{'='*60}", flush=True)
t0 = time.time()
# Initialize population: (POP, K) long tensors
pop = torch.zeros(POP, K, dtype=torch.long, device=DEVICE)
# Seed 0: greedy
g = greedy_dims[:K] if K <= len(greedy_dims) else greedy_dims + list(range(K - len(greedy_dims)))
pop[0] = torch.tensor(g[:K], device=DEVICE)
# Rest: random K-subsets of 768
for i in range(1, POP):
pop[i] = torch.randperm(768, device=DEVICE)[:K]
fits = batch_fitness(sub_f, sub_person, pop)
best_f1 = fits.max().item()
best_genome = pop[fits.argmax()].clone()
stag = 0
for gen in range(GEN):
# Sort
order = fits.argsort(descending=True)
pop = pop[order]
fits = fits[order]
if fits[0].item() > best_f1:
best_f1 = fits[0].item()
best_genome = pop[0].clone()
stag = 0
else:
stag += 1
# New population
new_pop = pop[:ELITE].clone()
# Immigration
n_imm = POP // 5 if stag > 200 else 0
if n_imm > 0:
imm = torch.stack([torch.randperm(768, device=DEVICE)[:K] for _ in range(n_imm)])
new_pop = torch.cat([new_pop, imm])
# Breed
n_breed = POP - new_pop.shape[0]
# Tournament selection
t1 = torch.randint(0, POP // 2, (n_breed, 5), device=DEVICE)
p1_idx = t1[torch.arange(n_breed, device=DEVICE), fits[t1].argmax(dim=1)]
t2 = torch.randint(0, POP // 2, (n_breed, 5), device=DEVICE)
p2_idx = t2[torch.arange(n_breed, device=DEVICE), fits[t2].argmax(dim=1)]
parents1 = pop[p1_idx] # (n_breed, K)
parents2 = pop[p2_idx]
# Crossover: for each position, pick from parent1 or parent2
mask = torch.rand(n_breed, K, device=DEVICE) < 0.5
children = torch.where(mask, parents1, parents2)
# Mutation: replace random positions with random dims
mut_rate = 0.05 * (1 + stag / 100)
mut_mask = torch.rand(n_breed, K, device=DEVICE) < mut_rate
random_dims = torch.randint(0, 768, (n_breed, K), device=DEVICE)
children = torch.where(mut_mask, random_dims, children)
new_pop = torch.cat([new_pop, children])[:POP]
pop = new_pop
fits = batch_fitness(sub_f, sub_person, pop)
if (gen + 1) % 100 == 0:
elapsed = time.time() - t0
gen_s = (gen + 1) / elapsed
print(f" gen {gen+1:5d}: best={fits.max().item():.4f} "
f"best_ever={best_f1:.4f} stag={stag} "
f"{gen_s:.0f} gen/s", flush=True)
if stag > 1000:
print(f" Converged at gen {gen+1}")
break
elapsed = time.time() - t0
best_dims = best_genome.cpu().tolist()
gates = K * 85
gens_done = gen + 1
print(f"\n WINNER: {K} dims, F1={best_f1:.4f}, {gates} gates, "
f"{elapsed:.1f}s, {gens_done/elapsed:.0f} gen/s", flush=True)
all_results.append({
"K": K, "best_f1": round(best_f1, 4), "genome": sorted(best_dims),
"gates": gates, "time_s": round(elapsed, 1),
"generations": gens_done, "gen_per_s": round(gens_done / elapsed),
})
print(f"\n{'='*60}")
print("Results:")
for r in all_results:
print(f" K={r['K']:3d} {r['gates']:6d} gates F1={r['best_f1']:.4f} "
f"{r['gen_per_s']} gen/s ({r['generations']} gen, {r['time_s']}s)")
out = os.path.join(SCRIPT_DIR, "evolved_extreme.json")
with open(out, "w") as f:
json.dump(all_results, f, indent=2)
print(f"Saved: {out}")
if __name__ == "__main__":
main()