import json, random import numpy as np import torch SEED = 2303 np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED) EPS = 1e-12 def normalize(x): x = np.maximum(np.asarray(x, dtype=float), EPS) return x / x.sum() def ait(p, q): p, q = normalize(p), normalize(q) D = len(p) lr = np.log(p[:, None] / p[None, :]) - np.log(q[:, None] / q[None, :]) return float((lr * lr).sum() / (2 * D)) def entropy(p): p = normalize(p) return float(-(p * np.log(p)).sum()) def hbin(s): return float(-s*np.log(s) - (1-s)*np.log(1-s)) def row(s, c): return np.r_[s, (1-s)*normalize(c)] def clr(c): z = np.log(normalize(c)); return z-z.mean() def checks(): # Prediction 1: fixed content gives an exact quadratic sink dependence. # Direct expansion of the supplied pairwise definition gives m/D*b^2, # not m/(2D)*b^2 as stated in the idea. m, D = 4, 5 c = normalize([.50, .30, .15, .05]) s = .80 sink_rows = [] for t in np.linspace(.08, .92, 18): b = np.log(s*(1-t)/(t*(1-s))) actual = ait(row(s,c), row(t,c)) sink_rows.append((float(b*b), actual, float(b))) x = np.array([v[0] for v in sink_rows]); y = np.array([v[1] for v in sink_rows]) coef = float(np.dot(x,y)/np.dot(x,x)) correct = m/D claimed = m/(2*D) # Prediction 2: the implemented content target is exactly quadratic in a # clr displacement, independent of sink mass. base = normalize([.50, .30, .15, .05]) content_rows = [] for a in np.linspace(0, 2.5, 12): c1 = normalize(np.exp(np.log(base) + a*clr(base))) c2 = normalize(np.exp(np.log(base) - a*clr(base))) implemented = float(np.sum((clr(c1)-clr(c2))**2)/m) expected = float((2*a)**2*np.sum(clr(base)**2)/m) content_rows.append((float(a), implemented, expected)) content_max_err = max(abs(r[1]-r[2]) for r in content_rows) # Prediction 3: entropy decomposition is exact. ent_err = [] for _ in range(20): z = np.random.dirichlet(np.ones(D)); s0 = z[0] ent_err.append(abs(entropy(z) - (hbin(s0) + (1-s0)*entropy(z[1:]/(1-s0))))) # Explicit counterexample to the claimed full decomposition: same content # makes the discrepancy especially transparent. s2, t2 = .82, .63 b2 = np.log(s2*(1-t2)/(t2*(1-s2))) actual2 = ait(row(s2,base), row(t2,base)) claimed2 = m/(2*D)*b2*b2 return {'prediction_1_sink_quadratic': { 'observed_coefficient': coef, 'predicted_from_definition_m_over_D': correct, 'paper_claim_m_over_2D': claimed, 'relative_error_correct': abs(coef-correct)/correct, 'relative_error_paper_claim': abs(coef-claimed)/claimed}, 'prediction_2_content_quadratic': { 'max_abs_error': content_max_err, 'mean_ratio_implemented_to_expected': float(np.mean([r[1]/r[2] for r in content_rows[1:]]))}, 'prediction_3_entropy_identity_max_abs_error': max(ent_err), 'decomposition_counterexample_same_content': { 'actual_pairwise_aitchison': actual2, 'paper_formula': claimed2, 'ratio_actual_to_paper': actual2/claimed2}, 'sink_sweep': sink_rows, 'content_sweep': content_rows} def toy_distill(): # Sink-heavy teacher and deliberately low-capacity student (one shared row). torch.manual_seed(SEED) R, M = 64, 7 sink = torch.full((R,1), .985) tc = torch.softmax(torch.randn(R,M)*1.1, dim=-1) teacher = torch.cat([sink, (1-sink)*tc], dim=1) def run(kind): sink_logit = torch.tensor(0., requires_grad=True) content_logits = torch.zeros(M, requires_grad=True) opt = torch.optim.SGD([sink_logit, content_logits], lr=.12) grad_max = 0. for _ in range(300): t = torch.sigmoid(sink_logit); c = torch.softmax(content_logits, dim=-1) student = torch.cat([t.expand(R,1), ((1-t)*c).expand(R,M)], dim=1) if kind == 'kl': loss = (teacher * (teacher.clamp_min(1e-12).log()-student.clamp_min(1e-12).log())).sum(1).mean() else: sinkloss = (torch.log(t/(1-t))-torch.log(teacher[:,0]/(1-teacher[:,0]))).square().mean() tz = torch.log(teacher[:,1:]/(1-teacher[:,0:1])); tz -= tz.mean(1, keepdim=True) sz = torch.log(c); sz -= sz.mean() contentloss = (tz-sz).square().sum(1).mean()/M loss = sinkloss + contentloss opt.zero_grad(); loss.backward() grad_max = max(grad_max, float(max(abs(sink_logit.grad.item()), content_logits.grad.abs().max().item()))) opt.step() with torch.no_grad(): t = torch.sigmoid(sink_logit); c = torch.softmax(content_logits,0) st = torch.cat([t.expand(R,1), ((1-t)*c).expand(R,M)],1) sink_err = float((st[:,0]-teacher[:,0]).abs().mean()) cz = torch.log(teacher[:,1:]/(1-teacher[:,0:1])); cz-=cz.mean(1,keepdim=True) sz = torch.log(c); sz-=sz.mean() cd = float(((cz-sz)**2).sum(1).mean()/M) kl = float((teacher*(teacher.log()-st.log())).sum(1).mean()) return {'final_kl':kl,'sink_abs_error':sink_err,'content_aitchison_squared':cd,'max_gradient':grad_max} return {'kl_baseline':run('kl'), 'sink_content_aitchison':run('idea')} def main(): out = {'seed': SEED, 'math_checks': checks(), 'toy_distillation': toy_distill()} with open('results.json','w') as f: json.dump(out,f,indent=2) print(json.dumps(out, indent=2)) if __name__ == '__main__': main()