kpshinnik's picture
download
raw
5.57 kB
"""
DANCE-lite: a faithful, compact realization of DANCE's hierarchical text
condensation + fusion classifier, on a text-attributed graph where each node's
"text" is its bag of present vocabulary words (each present word = one chunk).
Implements the trainable path of Sec 4.3-4.4:
g_v : 2-layer GCN graph embedding on BoW features
chunk gating : a_{v,c} = (W_s g_v) . E_c over the node's present words (Eq 8),
hard top-B_tok selection -> evidence embedding t~_v (Eq 9)
fusion : x_v = LN(W_g g_v + alpha_v W_t t~_v) (Eq 10)
classifier : Dec(x_v)
This is what the deletion/insertion evidence tests (Claim 3) probe, and what the
condensation accuracy/token comparison (Claim 1) trains.
"""
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
def normalized_adj(edge_index, n, device):
"""Symmetric normalized adjacency with self-loops (GCN)."""
row, col = edge_index
idx = torch.cat([edge_index, torch.arange(n, device=device).repeat(2, 1)], dim=1)
vals = torch.ones(idx.shape[1], device=device)
A = torch.sparse_coo_tensor(idx, vals, (n, n)).coalesce()
deg = torch.sparse.sum(A, dim=1).to_dense()
dinv = deg.pow(-0.5); dinv[torch.isinf(dinv)] = 0
r, c = A.indices()
v = A.values() * dinv[r] * dinv[c]
return torch.sparse_coo_tensor(A.indices(), v, (n, n)).coalesce()
class DanceLite(nn.Module):
def __init__(self, vocab, d, n_classes, hidden=64):
super().__init__()
self.W0 = nn.Linear(vocab, hidden)
self.W1 = nn.Linear(hidden, d)
self.E = nn.Embedding(vocab, d) # per-word (chunk) embedding
self.Ws = nn.Linear(d, d, bias=False) # chunk-attention query (Eq 8)
self.Wg = nn.Linear(d, d)
self.Wt = nn.Linear(d, d)
self.wgate = nn.Linear(2 * d, 1) # alpha_v gate (Eq 10)
self.ln = nn.LayerNorm(d)
self.lnt = nn.LayerNorm(d)
self.dec = nn.Linear(d, n_classes) # fused head (Eq 10-11)
self.dec_t = nn.Linear(d, n_classes) # text-evidence-only head (Dec_t, Eq 11)
self.d = d
def gcn(self, x, adj):
h = F.relu(self.W0(x))
h = torch.sparse.mm(adj, h)
h = self.W1(h)
h = torch.sparse.mm(adj, h)
return h # g_v : [n, d]
def set_padded(self, present_list, device):
"""Precompute a padded [n, Lmax] word-id tensor + validity mask once."""
n = len(present_list)
L = max((len(p) for p in present_list), default=1)
pad = torch.zeros(n, L, dtype=torch.long, device=device)
valid = torch.zeros(n, L, dtype=torch.bool, device=device)
for i, p in enumerate(present_list):
if p:
pad[i, :len(p)] = torch.as_tensor(p, device=device)
valid[i, :len(p)] = True
self._pad, self._valid, self._L = pad, valid, L
def evidence(self, g, present_list, B_tok, mask_override=None):
"""Vectorised over nodes. Attend over each node's present words (padded),
hard-select top-B_tok, return evidence embedding t~ [n,d] and packs.
mask_override[i]: optional iterable of allowed word ids (del/insertion)."""
n = g.shape[0]
if not hasattr(self, "_pad") or self._pad.shape[0] != n:
self.set_padded(present_list, g.device)
pad, valid = self._pad, self._valid.clone()
if mask_override is not None:
for i in range(n):
allowed = mask_override[i]
if len(allowed) == 0:
valid[i] = False; continue
al = torch.as_tensor(list(allowed), device=g.device)
keep = torch.isin(pad[i], al)
valid[i] = valid[i] & keep
q = self.Ws(g) # [n,d]
e = self.E(pad) # [n,L,d]
a = torch.einsum("nd,nld->nl", q, e) / np.sqrt(self.d) # Eq 8
a = a.masked_fill(~valid, float("-inf"))
w = torch.softmax(a, dim=1)
w = torch.nan_to_num(w, nan=0.0)
# hard top-B_tok per node
B = min(B_tok, self._L)
topv, topi = torch.topk(w, B, dim=1) # [n,B]
pi = torch.zeros_like(w).scatter(1, topi, topv)
denom = pi.sum(1, keepdim=True).clamp_min(1e-9)
pi = pi / denom
t = torch.einsum("nl,nld->nd", pi, e) # Eq 9
# packs: selected (word, weight) per node (only valid picks)
packs = []
pad_cpu = pad.detach().cpu().numpy()
topi_cpu = topi.detach().cpu().numpy()
pi_cpu = pi.detach().cpu().numpy()
valid_cpu = valid.detach().cpu().numpy()
for i in range(n):
pk = [(int(pad_cpu[i, j]), float(pi_cpu[i, j])) for j in topi_cpu[i]
if valid_cpu[i, j] and pi_cpu[i, j] > 0]
packs.append(pk)
return t, packs
def forward(self, x, adj, present_list, B_tok, mask_override=None):
g = self.gcn(x, adj)
t, packs = self.evidence(g, present_list, B_tok, mask_override)
alpha = torch.sigmoid(self.wgate(torch.cat([g, t], dim=1))) # Eq 10
xf = self.ln(self.Wg(g) + alpha * self.Wt(t))
logits_fused = self.dec(xf)
logits_text = self.dec_t(self.lnt(self.Wt(t))) # Dec_t on evidence only
return logits_fused, packs, g, logits_text
def present_words(x):
"""List of present-word indices per node from a binary BoW matrix."""
xn = (x > 0).cpu().numpy()
return [np.nonzero(row)[0].tolist() for row in xn]

Xet Storage Details

Size:
5.57 kB
·
Xet hash:
86106d7ef120cb2908330e0c42eadd21eb70d2831318a44d53a5e98c817a75df

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