Constant-sum ordinal preference loss / experiment.py

Mechanism failed

Raw ⬇ ZIP
  1import json, random
  2import numpy as np
  3import torch
  4from torch import nn
  5from scipy.stats import kendalltau
  6
  7SEED=17
  8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
  9torch.set_num_threads(4)
 10device='cuda' if torch.cuda.is_available() else 'cpu'
 11try:
 12    if device=='cuda': torch.zeros(1, device='cuda')
 13except Exception:
 14    device='cpu'
 15
 16def supplied_logprob(d, a, p):
 17    # Exact formula in the idea: log P(y)/P(y-1)=a[y]+p[y]d.
 18    inc=a[1:][None,:] + d[:,None]*p[1:][None,:]
 19    lr=torch.cat([torch.zeros((len(d),1),device=d.device), torch.cumsum(inc,1)],1)
 20    return torch.log_softmax(lr,1)
 21
 22def paper_ac_logprob(d, alpha, points):
 23    # Paper-faithful AC model: category logits alpha[y]+points[y]*d.
 24    return torch.log_softmax(alpha[None,:] + d[:,None]*points[None,:],1)
 25
 26def math_checks():
 27    # Verify recursion and test the actual sufficiency signature: for equal-N,
 28    # equal-score records, log-likelihood differences may depend on intercepts,
 29    # but must not depend on the skill difference d.
 30    a=torch.tensor([0.,-.2,.35],dtype=torch.float64)
 31    p=torch.tensor([0.,1.,2.],dtype=torch.float64)  # p0+p2=2
 32    ds=[-.8,.9]
 33    lp=[supplied_logprob(torch.tensor([d],dtype=torch.float64),a,p)[0] for d in ds]
 34    err=float(max((z[1:]-z[:-1]-(a[1:]+ds[k]*p[1:])).abs().max() for k,z in enumerate(lp)))
 35    c1=torch.tensor([2.,0.,1.]); c2=torch.tensor([1.,2.,0.])
 36    supplied_gap=[float(abs(sum((c1-c2)[y]*lp[k][y] for y in range(3)))) for k in range(2)]
 37    q=torch.tensor([0.,.5,1.],dtype=torch.float64)
 38    alpha=torch.tensor([0.,-.2,.35],dtype=torch.float64)
 39    ac=[paper_ac_logprob(torch.tensor([d],dtype=torch.float64),alpha,q)[0] for d in ds]
 40    ac_gap=[float(abs(sum((c1-c2)[y]*ac[k][y] for y in range(3)))) for k in range(2)]
 41    # For the supplied recursion the effective category slopes are cumulative
 42    # [0,p1,p1+p2], not p itself.
 43    return {'supplied_adjacent_error':err,'constant_sum_p0_plus_p2':float(p[0]+p[2]),
 44            'supplied_equal_score_loglik_gaps_at_d':supplied_gap,
 45            'paper_AC_equal_score_loglik_gaps_at_d':ac_gap,
 46            'supplied_gap_change':abs(supplied_gap[1]-supplied_gap[0]),
 47            'paper_AC_gap_change':abs(ac_gap[1]-ac_gap[0]),
 48            'supplied_effective_category_slopes':[0.,float(p[1]),float(p[1]+p[2])]}
 49
 50class Net(nn.Module):
 51    def __init__(self):
 52        super().__init__(); self.f=nn.Sequential(nn.Linear(5,16),nn.Tanh(),nn.Linear(16,1))
 53    def diff(self,x1,x2): return (self.f(x1)-self.f(x2)).squeeze(-1)
 54
 55class Supplied(nn.Module):
 56    def __init__(self):
 57        super().__init__(); self.net=Net(); self.a=nn.Parameter(torch.zeros(3)); self.v=nn.Parameter(torch.tensor(.5))
 58    def probs(self,x1,x2):
 59        d=self.net.diff(x1,x2); p=torch.stack([2-self.v,torch.ones((),device=d.device),self.v])
 60        return supplied_logprob(d,self.a,p).exp()
 61    def forward(self,x1,x2,y):
 62        pr=self.probs(x1,x2); return -pr[torch.arange(len(y),device=y.device),y].clamp_min(1e-8).log().mean()
 63
 64class PaperAC(nn.Module):
 65    def __init__(self):
 66        super().__init__(); self.net=Net(); self.alpha=nn.Parameter(torch.zeros(3)); self.v=nn.Parameter(torch.tensor(.5))
 67    def probs(self,x1,x2):
 68        d=self.net.diff(x1,x2); points=torch.stack([torch.zeros((),device=d.device),self.v,torch.ones((),device=d.device)])
 69        return paper_ac_logprob(d,self.alpha,points).exp()
 70    def forward(self,x1,x2,y):
 71        pr=self.probs(x1,x2); return -pr[torch.arange(len(y),device=y.device),y].clamp_min(1e-8).log().mean()
 72
 73class CE(nn.Module):
 74    def __init__(self):
 75        super().__init__(); self.net=Net(); self.head=nn.Linear(1,3)
 76    def probs(self,x1,x2): return self.head(self.net.diff(x1,x2)[:,None]).softmax(1)
 77    def forward(self,x1,x2,y): return nn.functional.cross_entropy(self.head(self.net.diff(x1,x2)[:,None]),y)
 78
 79def data(n=2200,items=35):
 80    x=torch.randn(items,5); w=torch.tensor([1.1,-.7,.4,.9,-.5]); skill=x@w
 81    i=torch.randint(items,(n,)); j=torch.randint(items,(n,)); j[i==j]=(j[i==j]+1)%items
 82    d=skill[i]-skill[j]; points=torch.tensor([0.,.5,1.]); alpha=torch.tensor([0.,-.25,0.])
 83    y=torch.distributions.Categorical(logits=alpha[None,:]+points[None,:]*d[:,None]).sample()
 84    return x[i],x[j],y,i,j,skill
 85
 86def run(cls, seed):
 87    torch.manual_seed(seed); tr=data(); te=data(); model=cls().to(device); opt=torch.optim.Adam(model.parameters(),lr=.025)
 88    x1,x2,y,*_=tr; x1,x2,y=x1.to(device),x2.to(device),y.to(device)
 89    for _ in range(100):
 90        opt.zero_grad(); loss=model(x1,x2,y); loss.backward(); opt.step()
 91    a,b,yy,i,j,true=te
 92    with torch.no_grad():
 93        pr=model.probs(a.to(device),b.to(device)).cpu(); nll=float(-pr[torch.arange(len(yy)),yy].clamp_min(1e-8).log().mean()); acc=float((pr.argmax(1)==yy).float().mean())
 94        # Recover item ranking by scorer output, which is the intended neural utility.
 95        scores=model.net.f(torch.randn(1,5).to(device)) if False else model.net.f(torch.eye(5,device=device)[:1])
 96        pred_score=model.net.f(torch.randn(1,5).to(device)) if False else model.net.f(te[0].to(device)).squeeze()
 97        # Pairwise held-out ranking concordance against known latent ordering.
 98        pred_d=(model.net.f(a.to(device))-model.net.f(b.to(device))).squeeze().cpu().numpy(); true_d=(true[i]-true[j]).numpy()
 99        tau=float(kendalltau(pred_d,true_d).statistic)
100    return nll,acc,tau
101
102def main():
103    out={'device':device,'math_checks':math_checks()}
104    for name,cls in [('supplied_ordinal',Supplied),('paper_constant_sum_AC',PaperAC),('multiclass_CE',CE)]:
105        vals=[run(cls,s) for s in (21,22,23)]
106        out[name]={'nll_mean':float(np.mean([v[0] for v in vals])),'accuracy_mean':float(np.mean([v[1] for v in vals])),'pair_rank_tau_mean':float(np.mean([v[2] for v in vals])),'per_seed':vals}
107    print(json.dumps(out,indent=2))
108if __name__=='__main__': main()