kpshinnik's picture
download
raw
9.71 kB
"""
Claim 2 (architecture invariants) + Claim 3 (evidence packs, deletion/insertion)
on real Cora (a text-attributed graph; each present vocabulary word = one text
chunk/span). Federated training and LLM surrogates are out of local scope (HF
Jobs credit-blocked); this is a single-client, BoW-text realization of DANCE's
mechanisms — enough to check the structural claims and evidence faithfulness.
"""
import json, os, sys, time
import numpy as np
import torch
sys.path.insert(0, os.path.dirname(__file__))
from dance_core import (label_aware_node_condensation, budgeted_neighbor_gating,
budgeted_chunk_selection, self_expressive_topology)
from dance_model import DanceLite, normalized_adj, present_words
torch.manual_seed(0); np.random.seed(0)
DEV = "cpu"
def load_cora():
from torch_geometric.datasets import Planetoid
ds = Planetoid(root='outputs/data/Cora', name='Cora')
return ds[0]
def hop_neighbors(edge_index, n, v, max_hop=2):
from collections import deque
adj = [[] for _ in range(n)]
ei = edge_index.numpy()
for a, b in zip(ei[0], ei[1]):
adj[a].append(b)
seen = {v: 0}; dq = deque([v]); hops = {0: [v], 1: [], 2: []}
while dq:
u = dq.popleft(); h = seen[u]
if h == max_hop: continue
for w in adj[u]:
if w not in seen:
seen[w] = h + 1; hops[h + 1].append(w); dq.append(w)
return hops
def claim2_invariants(data):
n = data.num_nodes
Z = data.x.numpy().astype(np.float32)
y = data.y.numpy()
r = 0.08
# --- node condensation: K = ceil(r n), label-stratified, every 10 rounds ---
core, _, quotas = label_aware_node_condensation(Z, y, r=r, P=8, seed=0)
K = int(np.ceil(r * n))
inv = {}
inv["n"] = n
inv["condensation_ratio"] = r
inv["K_expected"] = K
inv["K_actual"] = len(core)
inv["K_matches_ceil_rn"] = len(core) == K
# cadence: condensation runs at rounds 0,10,20,... ; between refreshes reuse
rounds = 30
refresh_rounds = [t for t in range(rounds) if t % 10 == 0]
inv["refresh_rounds_in_30"] = refresh_rounds
inv["refresh_every_10"] = refresh_rounds == [0, 10, 20]
# label stratification: per-class core proportion ~ per-class quota
core_labels = y[core]
inv["classes_covered"] = int(len(np.unique(core_labels)))
inv["total_classes"] = int(len(np.unique(y)))
# --- neighbor gating: budgets enforced across 0/1/2-hop + chunk budget ---
d = 32
rng = np.random.default_rng(0)
g = rng.normal(size=(n, d)).astype(np.float32)
t = rng.normal(size=(n, d)).astype(np.float32)
Wq = rng.normal(size=(d, d)).astype(np.float32) * 0.1
Wk = rng.normal(size=(d, d)).astype(np.float32) * 0.1
budgets = {0: 1, 1: 5, 2: 5}
B_tok = 8
viol_hop = 0; viol_tok = 0; checked = 0
for v in core[:50]:
hn = hop_neighbors(data.edge_index, n, int(v))
sel = budgeted_neighbor_gating(int(v), g, t, hn, Wq, Wk, budgets)
for ell in [0, 1, 2]:
if len(sel[ell]) > budgets[ell]:
viol_hop += 1
# chunk selection over selected neighbors' chunks (simulate 4 chunks/node)
cand_nodes = [u for ell in sel for u in sel[ell]]
C = len(cand_nodes) * 4 + 1
chunk_owner = np.repeat(cand_nodes, 4) if cand_nodes else np.array([int(v)])
chunk_embeds = rng.normal(size=(len(chunk_owner), d)).astype(np.float32)
Ws = rng.normal(size=(d, d)).astype(np.float32) * 0.1
csel, _ = budgeted_chunk_selection(int(v), g, chunk_embeds, chunk_owner,
cand_nodes or [int(v)], Ws, B_tok)
if len(csel) > B_tok:
viol_tok += 1
checked += 1
inv["gating_nodes_checked"] = checked
inv["hop_budget_violations"] = viol_hop
inv["chunk_budget_violations"] = viol_tok
inv["max_neighbors_per_core_node"] = sum(budgets.values())
inv["budgets"] = budgets
inv["B_tok"] = B_tok
# --- topology: O(Kq) coefficients, O(Kk) edges ---
Xcore = Z[core]
Sprior = (Xcore @ Xcore.T)
Sprior = (Sprior - Sprior.min()) / (Sprior.max() - Sprior.min() + 1e-9)
q, k = 10, 6
A, Zc, mask = self_expressive_topology(Xcore, Sprior, q=q, k=k, steps=40)
inv["topology_q"] = q; inv["topology_k"] = k
inv["Z_nnz"] = int((Zc != 0).sum())
inv["Z_nnz_bound_Kq"] = 2 * K * q
inv["Z_nnz_within_O_Kq"] = int((Zc != 0).sum()) <= 2 * K * q
inv["edges"] = int((A > 0).sum())
inv["edges_bound_Kk"] = 2 * K * k
inv["edges_within_O_Kk"] = int((A > 0).sum()) <= 2 * K * k
inv["dense_K2"] = K * K
inv["all_invariants_pass"] = bool(
inv["K_matches_ceil_rn"] and inv["refresh_every_10"] and
viol_hop == 0 and viol_tok == 0 and
inv["Z_nnz_within_O_Kq"] and inv["edges_within_O_Kk"])
return inv
def train_dance(data, B_tok=8, epochs=80):
n = data.num_nodes
vocab = data.x.shape[1]
d = 64
ncls = int(data.y.max()) + 1
x = (data.x > 0).float().to(DEV)
adj = normalized_adj(data.edge_index.to(DEV), n, DEV)
pres = present_words(data.x)
model = DanceLite(vocab, d, ncls).to(DEV)
opt = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4)
y = data.y.to(DEV)
tr = data.train_mask.to(DEV); te = data.test_mask.to(DEV)
for ep in range(epochs):
model.train(); opt.zero_grad()
logits, _, _, logits_t = model(x, adj, pres, B_tok)
loss = (torch.nn.functional.cross_entropy(logits[tr], y[tr])
+ torch.nn.functional.cross_entropy(logits_t[tr], y[tr]))
loss.backward(); opt.step()
model.eval()
with torch.no_grad():
logits, packs, _, logits_t = model(x, adj, pres, B_tok)
acc = float((logits[te].argmax(1) == y[te]).float().mean())
acc_t = float((logits_t[te].argmax(1) == y[te]).float().mean())
return model, x, adj, pres, y, te, acc, acc_t, packs
@torch.no_grad()
def eval_masked(model, x, adj, pres, y, te, B_tok, mask_override, base_pred):
"""Return (fused acc, text-head acc, text-head confidence for base predicted class)."""
logits, _, _, logits_t = model(x, adj, pres, B_tok, mask_override=mask_override)
acc_f = float((logits[te].argmax(1) == y[te]).float().mean())
acc_t = float((logits_t[te].argmax(1) == y[te]).float().mean())
prob_t = torch.softmax(logits_t, dim=1)
conf = float(prob_t[te].gather(1, base_pred[te].unsqueeze(1)).mean())
return acc_f, acc_t, conf
def claim3_evidence(data):
(model, x, adj, pres, y, te, base_acc, base_acc_t, _) = train_dance(data, B_tok=8, epochs=80)
n = data.num_nodes
rng = np.random.default_rng(0)
with torch.no_grad():
logits, _, _, logits_t = model(x, adj, pres, 8)
base_pred = logits_t.argmax(1) # text-head predicted class (what evidence must support)
ks = [1, 2, 4, 8, 16]
out = {"base_accuracy_fused": round(base_acc, 4),
"base_accuracy_texthead": round(base_acc_t, 4), "ks": ks}
ins_sel_c, ins_rand_c, del_sel_c, del_rand_c = [], [], [], []
ins_sel_a, ins_rand_a = [], []
full = [set(p) for p in pres]
for k in ks:
with torch.no_grad():
_, packs, _, _ = model(x, adj, pres, k)
sel = [set(w for w, _ in p) for p in packs]
rnd = [set(rng.choice(pres[i], size=min(k, len(pres[i])), replace=False).tolist())
if len(pres[i]) else set() for i in range(n)]
_, a_is, c_is = eval_masked(model, x, adj, pres, y, te, k, sel, base_pred)
_, a_ir, c_ir = eval_masked(model, x, adj, pres, y, te, k, rnd, base_pred)
_, _, c_ds = eval_masked(model, x, adj, pres, y, te, 9999,
[full[i] - sel[i] for i in range(n)], base_pred)
_, _, c_dr = eval_masked(model, x, adj, pres, y, te, 9999,
[full[i] - rnd[i] for i in range(n)], base_pred)
ins_sel_a.append(a_is); ins_rand_a.append(a_ir)
ins_sel_c.append(c_is); ins_rand_c.append(c_ir)
del_sel_c.append(c_ds); del_rand_c.append(c_dr)
out.update({
"insertion_selected_acc": [round(a, 4) for a in ins_sel_a],
"insertion_random_acc": [round(a, 4) for a in ins_rand_a],
"insertion_selected_conf": [round(a, 4) for a in ins_sel_c],
"insertion_random_conf": [round(a, 4) for a in ins_rand_c],
"deletion_selected_conf": [round(a, 4) for a in del_sel_c],
"deletion_random_conf": [round(a, 4) for a in del_rand_c],
# sufficiency: keeping only selected evidence recovers text-head confidence
"sufficiency_gap_conf": round(float(np.mean(ins_sel_c) - np.mean(ins_rand_c)), 4),
# necessity: removing selected evidence drops confidence more than random
"necessity_gap_conf": round(float(np.mean(del_rand_c) - np.mean(del_sel_c)), 4),
"sufficient": bool(np.mean(ins_sel_c) > np.mean(ins_rand_c) + 0.02),
"necessary": bool(np.mean(del_sel_c) < np.mean(del_rand_c) - 0.02),
})
return out
def main():
data = load_cora()
print("Cora:", data.num_nodes, "nodes")
t0 = time.time()
inv = claim2_invariants(data)
print("Claim 2 invariants:", inv["all_invariants_pass"], f"({time.time()-t0:.1f}s)")
t0 = time.time()
ev = claim3_evidence(data)
print("Claim 3 evidence: sufficient=", ev["sufficient"], "necessary=", ev["necessary"],
f"({time.time()-t0:.1f}s)")
out = {"dataset": "Cora", "claim2_invariants": inv, "claim3_evidence": ev}
os.makedirs("outputs", exist_ok=True)
with open("outputs/claim23_cora.json", "w") as f:
json.dump(out, f, indent=2)
print(json.dumps(out, indent=2))
if __name__ == "__main__":
main()

Xet Storage Details

Size:
9.71 kB
·
Xet hash:
62e596ee1b5c1f1d643f104626e5e915374bd0c5612a9a32bd6746aab693c54c

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.