"""C4 exact small-system checks and symbolic elimination. Python standard library only."""
from fractions import Fraction as F
from itertools import combinations
from math import sin,cos,pi,sqrt
from pathlib import Path
import json,sys
OUT=Path(sys.argv[1]) if len(sys.argv)>1 else Path(__file__).resolve().parents[1]/'qa/sparse/verification-results.json'
def clean(x):
 if isinstance(x,F):return int(x) if x.denominator==1 else str(x)
 if isinstance(x,dict):return {str(k):clean(v) for k,v in x.items()}
 if isinstance(x,(list,tuple)):return [clean(v) for v in x]
 return x
def mat(n,m=None):return [[F(0) for _ in range(n if m is None else m)] for _ in range(n)]
def eye(n):return [[F(i==j) for j in range(n)] for i in range(n)]
def tr(A):return list(map(list,zip(*A)))
def mv(A,x):return [sum(a*v for a,v in zip(row,x)) for row in A]
def mm(A,B):return [[sum(a*b for a,b in zip(row,col)) for col in zip(*B)] for row in A]
def add(A,B,s=1):return [[a+s*b for a,b in zip(ar,br)] for ar,br in zip(A,B)]
def va(a,b,s=1):return [x+s*y for x,y in zip(a,b)]
def dot(a,b):return sum(x*y for x,y in zip(a,b))
def energy(A,x):return dot(x,mv(A,x))
def scale(A,t):return [[t*x for x in row] for row in A]
def solve(A,b):
 n=len(A);W=[list(map(F,row))+[F(v)] for row,v in zip(A,b)]
 for k in range(n):
  q=next(i for i in range(k,n) if W[i][k]);W[k],W[q]=W[q],W[k];t=W[k][k];W[k]=[v/t for v in W[k]]
  for i in range(n):
   if i!=k:
    t=W[i][k];W[i]=[v-t*w for v,w in zip(W[i],W[k])]
 return [row[-1] for row in W]
def inv(A):return tr([solve(A,e) for e in eye(len(A))])
def sub(A,I,J):return [[A[i][j] for j in J] for i in I]
def chain(n):return [[F(2 if i==j else -1 if abs(i-j)==1 else 0) for j in range(n)] for i in range(n)]
def interp(n):
 assert n>=3 and n%2==1;P=mat(n,(n-1)//2)
 for j in range((n-1)//2):P[2*j][j]=F(1,2);P[2*j+1][j]=1;P[2*j+2][j]=F(1,2)
 return P
def symbolic(n,edges,order):
 assert sorted(order)==list(range(n));pos={v:i for i,v in enumerate(order)};g=[set() for _ in range(n)]
 for u,v in edges:g[u].add(v);g[v].add(u)
 states=[];cols=[]
 for k,v in enumerate(order):
  N=sorted(g[v],key=pos.get);fill=[]
  for a,b in combinations(N,2):
   if b not in g[a]:g[a].add(b);g[b].add(a);fill.append([a,b])
  cols.append([k]+[pos[w] for w in N]);states.append(dict(vertex=v,later_neighbors=N,new_fill=fill))
  for w in N:g[w].remove(v)
  g[v].clear()
 d=[len(x)-1 for x in cols]
 return dict(order=order,states=states,columns=cols,parent=[c[1] if len(c)>1 else None for c in cols],nnz_L=sum(map(len,cols)),fill_count=sum(len(x['new_fill']) for x in states),symmetric_updates=sum(t*(t+1)//2 for t in d),sum_column_squares=sum((t+1)**2 for t in d))
def grid(q):return [(q*i+j,q*(i+di)+j+dj) for i in range(q) for j in range(q) for di,dj in [(1,0),(0,1)] if i+di<q and j+dj<q]
def nd(q):
 layers=[]
 def rec(r0,r1,c0,c1,depth):
  if r0==r1 or c0==c1:return []
  if r1-r0==1 and c1-c0==1:return [q*r0+c0]
  rm=(r0+r1)//2;cm=(c0+c1)//2
  S=[q*i+j for i in range(r0,r1) for j in range(c0,c1) if i==rm or j==cm];layers.append(dict(depth=depth,separator=S))
  return sum([rec(a,b,c,d,depth+1) for a,b in [(r0,rm),(rm+1,r1)] for c,d in [(c0,cm),(cm+1,c1)]],[])+S
 return rec(0,q,0,q,0),layers
def ldl(A,pattern=None):
 A=[list(map(F,row)) for row in A];n=len(A);L=eye(n);D=[]
 for j in range(n):
  d=A[j][j]-sum(L[j][k]**2*D[k] for k in range(j));D.append(d)
  if d<=0:return dict(status='nonpositive_pivot',index=j,L=L,D=D)
  for i in range(j+1,n):
   if pattern is None or (i,j) in pattern:L[i][j]=(A[i][j]-sum(L[i][k]*D[k]*L[j][k] for k in range(j)))/d
 return dict(status='ok',L=L,D=D)
def ldl_matrix(d):return mm([[x*d['D'][j] for j,x in enumerate(row)] for row in d['L']],tr(d['L']))
def ldl_apply(d,r):
 L=d['L'];n=len(r);y=list(r)
 for i in range(n):y[i]-=sum(L[i][j]*y[j] for j in range(i))
 z=[v/t for v,t in zip(y,d['D'])]
 for i in range(n-1,-1,-1):z[i]-=sum(L[j][i]*z[j] for j in range(i+1,n))
 return z
def pcg(A,b,apply):
 n=len(b);x=[F(0)]*n;r=list(b);z=apply(r);p=z[:];rho=dot(r,z);hist=[dict(x=x[:],residual=r[:],norm2=dot(r,r))]
 for k in range(n):
  Ap=mv(A,p);alpha=rho/dot(p,Ap);x=va(x,p,alpha);r=va(r,Ap,-alpha);hist.append(dict(x=x[:],residual=r[:],norm2=dot(r,r)))
  assert r==va(b,mv(A,x),-1)
  if not any(r):return hist
  z=apply(r);n_rho=dot(r,z);p=va(z,p,n_rho/rho);rho=n_rho
 raise AssertionError('PCG failed exact n-step termination')
def smoother(A):return [[F(2,3)*F(i==j)/A[i][i] for j in range(len(A))] for i in range(len(A))]
def hierarchy(A):
 As=[A];Ps=[]
 while len(A)>1:
  P=interp(len(A));A=mm(mm(tr(P),A),P);Ps.append(P);As.append(A)
 return As,Ps
def vcycle(As,Ps,l,b,trace=None):
 A=As[l];n=len(A)
 if n==1:
  z=solve(A,b)
  if trace is not None:trace.append(dict(level=l,n=n,rhs=b,coarse_exact=z))
  return z
 W=smoother(A);x=mv(W,b);r=va(b,mv(A,x),-1);rc=mv(tr(Ps[l]),r)
 if trace is not None:trace.append(dict(level=l,n=n,rhs=b,pre=x[:],residual=r,coarse_rhs=rc))
 zc=vcycle(As,Ps,l+1,rc,trace);corr=mv(Ps[l],zc);x=va(x,corr);before=x[:];x=va(x,mv(W,va(b,mv(A,x),-1)))
 if trace is not None:trace.append(dict(level=l,n=n,coarse_solution=zc,prolongated=corr,corrected=before,post=x[:],residual_after=va(b,mv(A,x),-1)))
 return x
def error_matrix(As,Ps,l):
 A=As[l];n=len(A)
 if n==1:return mat(1)
 Echild=error_matrix(As,Ps,l+1);Bc=mm(add(eye(len(As[l+1])),Echild,-1),inv(As[l+1]));P=Ps[l];W=smoother(A);S=add(eye(n),mm(W,A),-1)
 return mm(mm(S,add(eye(n),mm(mm(mm(P,Bc),tr(P)),A),-1)),S)
def run():
 R={};checks=0
 for mask in range(1<<10):
  edges0=[e for j,e in enumerate(combinations(range(5),2)) if mask>>j&1]
  for order0 in [list(range(5)),list(reversed(range(5)))]:
   sym=symbolic(5,edges0,order0);pos={v:i for i,v in enumerate(order0)};g=[set() for _ in range(5)]
   for u,v in edges0:g[pos[u]].add(pos[v]);g[pos[v]].add(pos[u])
   for j in range(5):
    for i in range(j+1,5):
     seen={j};todo=[j]
     while todo:
      v=todo.pop()
      for w in g[v]:
       if w not in seen and (w<j or w==i):seen.add(w);todo.append(w)
     assert (i in sym['columns'][j])==(i in seen)
     if i in sym['columns'][j]:
      v=j
      while v is not None and v<i:v=sym['parent'][v]
      assert v==i
   checks+=1
 R['structural_exhaustive_checks']=dict(graphs=1024,orderings_per_graph=2,path_and_ancestor_checks=checks)
 cancel=[[F(v) for v in row] for row in [[1,1,1],[1,2,1],[1,1,2]]];cf=ldl(cancel);assert cf['status']=='ok' and cf['L'][2][1]==0
 R['numerical_cancellation']=dict(A=cancel,exact_LDL=cf,symbolic=symbolic(3,[(0,1),(0,2),(1,2)],[0,1,2]))
 A=chain(5);I=[1,3];B=[0,2,4];Q=inv(sub(A,I,I));S=add(sub(A,B,B),mm(mm(sub(A,B,I),Q),sub(A,I,B)),-1);rhs=va([F(1)]*3,mv(mm(sub(A,B,I),Q),[F(1)]*2),-1);xb=solve(S,rhs);xi=mv(Q,va([F(1)]*2,mv(sub(A,I,B),xb),-1));x=[F(0)]*5
 for i,v in zip(I,xi):x[i]=v
 for i,v in zip(B,xb):x[i]=v
 assert mv(A,x)==[1]*5
 R['schur']=dict(A=A,interior=I,boundary=B,S=S,reduced_rhs=rhs,x=x,wrong_deleted_solution=solve(sub(A,B,B),[1]*3))
 edges=[(0,1),(1,2),(2,3),(3,4),(4,5),(5,0),(1,4)]
 R['symbolic']=[symbolic(6,edges,list(range(6))),symbolic(6,edges,[0,2,5,3,1,4])]
 # Ancestor is necessary, not sufficient: path0-1-2 plus3-4-5, joined at6.
 es=[(0,1),(1,2),(2,6),(3,4),(4,5),(5,6)]
 R['elimination_tree']=symbolic(7,es,list(range(7)))
 assert 2 not in R['elimination_tree']['columns'][0] and R['elimination_tree']['parent'][0]==1 and R['elimination_tree']['parent'][1]==2
 R['nested_dissection']=[]
 for q in [3,7,15,31]:
  order,layers=nd(q);a=symbolic(q*q,grid(q),list(range(q*q)));b=symbolic(q*q,grid(q),order)
  R['nested_dissection'].append(dict(q=q,natural={k:a[k] for k in ['nnz_L','fill_count','sum_column_squares']},nested={k:b[k] for k in ['nnz_L','fill_count','sum_column_squares']},**(dict(order=order,layers=layers,columns=b['columns']) if q==7 else {})))
 A=mat(9)
 for i in range(9):A[i][i]=4
 for i,j in grid(3):A[i][j]=A[j][i]=-1
 pattern={(i,j) for i in range(9) for j in range(i) if A[i][j]};IC=ldl(A,pattern);assert IC['status']=='ok';M=ldl_matrix(IC);res=add(A,M,-1)
 assert all(res[i][j]==0 for i in range(9) for j in range(9) if i==j or A[i][j])
 b=[F(i%3+1) for i in range(9)];h0=pcg(A,b,lambda r:r[:]);hi=pcg(A,b,lambda r:ldl_apply(IC,r))
 bad=[[F(v,5) for v in row] for row in [[5,3,0,3],[3,5,3,0],[0,3,5,-3],[3,0,-3,5]]];p={(i,j) for i in range(4) for j in range(i) if bad[i][j]};full=ldl(bad);broken=ldl(bad,p);fixed=ldl(add(bad,scale(eye(4),F(1,5))),p)
 assert full['status']=='ok' and broken['D'][-1]==F(-32,175) and fixed['status']=='ok'
 R['incomplete_cholesky']=dict(A=A,factor=IC,A_minus_M=res,cg=h0,pcg=hi,breakdown=dict(A=bad,exact=full,incomplete=broken,shift=F(1,5),shifted=fixed))
 A=chain(9);subs=[list(range(6)),list(range(3,9))];B=mat(9);local=[];r=[F(1)]*9
 for ids in subs:
  Ai=sub(A,ids,ids);zi=solve(Ai,[r[i] for i in ids]);local.append(dict(ids=ids,solution=zi));Ji=inv(Ai)
  for ii,i in enumerate(ids):
   for jj,j in enumerate(ids):B[i][j]+=Ji[ii][jj]
 p=list(map(F,[1,2,3,4,5,4,3,2,1]));B0=scale(mm([[v] for v in p],[p]),1/energy(A,p));B2=add(B,B0);e1=va(p,mv(B,mv(A,p)),F(-1,2));e2=va(p,mv(B2,mv(A,p)),F(-1,2))
 assert ldl(B)['status']=='ok' and ldl(B2)['status']=='ok'
 R['schwarz']=dict(A=A,subdomains=subs,unit_rhs_local=local,sum_correction=mv(B,r),coarse_basis=p,coarse_A=energy(A,p),slow_rhs=mv(A,p),half_step_one=e1,half_step_two=e2,energies=[energy(A,p),energy(A,e1),energy(A,e2)],coarse_overcount=va(p,mv(B2,mv(A,p)),-1))
 n=15;mu=lambda k,w:1-w*(1-cos(k*pi/(n+1)));e=[sin(pi*i/16)+sin(12*pi*i/16) for i in range(1,16)];sm=mv([[float(v) for v in row] for row in add(eye(n),mm(smoother(chain(n)),chain(n)),-1)],e)
 R['smoothing']=dict(n=n,omega=F(2,3),mu1=mu(1,2/3),mu12=mu(12,2/3),unweighted_mu15=mu(15,1),error=e,after=sm)
 A=chain(7);P=interp(7);Ac=mm(mm(tr(P),A),P);B=mm(mm(P,inv(Ac)),tr(P));C=add(eye(7),mm(B,A),-1);z=[F(1),F(2),F(1)];ep=mv(P,z);e=list(map(F,[-1,1,-1,1,-1,1,-1]));rc=mv(tr(P),mv(A,e));ec=solve(Ac,rc);after=mv(C,e)
 assert mv(C,ep)==[0]*7 and rc==[F(1,2),0,F(1,2)] and mm(C,C)==C and mm(tr(P),mm(A,C))==mat(3,7)
 invisible=[F(1),0,0,0,0,0,0];iv=mv(C,invisible);assert mv(tr(P),mv(A,iv))==[0]*3
 R['two_grid']=dict(A=A,P=P,Ac=Ac,coarse_error=ep,alternating=e,restricted_Ae=rc,coarse_solution=ec,after=after,energy_before=energy(A,e),energy_after=energy(A,after),invisible_example=iv)
 R['vcycle']=[]
 for n in [7,15]:
  As,Ps=hierarchy(chain(n));b=mv(As[0],[F(1)]*n);trace=[];x=vcycle(As,Ps,0,b,trace);e=va([F(1)]*n,x,-1);BA=tr([vcycle(As,Ps,0,col) for col in eye(n)])
  assert BA==tr(BA) and ldl(BA)['status']=='ok'
  E=error_matrix(As,Ps,0);assert add(eye(n),mm(BA,As[0]),-1)==E and mv(E,[F(1)]*n)==e
  hist=pcg(As[0],b,lambda v:vcycle(As,Ps,0,v))
  R['vcycle'].append(dict(n=n,As=As,Ps=Ps,b=b,trace=trace,solution=x,error=e,residual=va(b,mv(As[0],x),-1),energy=energy(As[0],e),initial_energy=F(2),preconditioner_spd=True,error_propagation_matrix_equal=True,pcg_steps=len(hist)-1))
 A=chain(8);P0=[[F(i//2==j) for j in range(4)] for i in range(8)];J=add(eye(8),mm(smoother(A),A),-1);P=mm(J,P0);Ac=mm(mm(tr(P),A),P);ones=[F(1)]*8;repro=mv(P,[F(1)]*4)
 assert repro==[F(2,3)]+[F(1)]*6+[F(2,3)] and ldl(Ac)['status']=='ok'
 C0=add(eye(8),mm(mm(mm(P0,inv(mm(mm(tr(P0),A),P0))),tr(P0)),A),-1);C=add(eye(8),mm(mm(mm(P,inv(Ac)),tr(P)),A),-1);r0=mv(C0,ones);rs=mv(C,ones)
 R['smoothed_aggregation']=dict(A=A,P0=P0,P=P,Ac=Ac,constant_reproduction=repro,defect=va(ones,repro,-1),Aones=mv(A,ones),energy_columns_before=[energy(A,c) for c in tr(P0)],energy_columns_after=[energy(A,c) for c in tr(P)],nnz_A=sum(bool(v) for row in A for v in row),nnz_P0=sum(bool(v) for row in P0 for v in row),nnz_P=sum(bool(v) for row in P for v in row),nnz_Ac=sum(bool(v) for row in Ac for v in row),constant_errors=[r0,rs],constant_energies=[energy(A,r0),energy(A,rs)])
 OUT.parent.mkdir(parents=True,exist_ok=True);OUT.write_text(json.dumps(clean(R),ensure_ascii=False,indent=2)+'\n');print('All C4 exact and structural checks passed',OUT)
if __name__=='__main__':run()
