#!/usr/bin/env python3
"""Hand-designed MLP and dictionary constructions from Lesson 2.3/Lab 07."""
import json,math,platform
from numbers import Real
WIN=((1,1,0),(0,-1,1),(1,0,-1),(-1,1,1));BIN=(0,-1,1,1)
WOUT=((1,0,-1,1),(0,1,0,-2),(1,-1,2,0));BOUT=(0,.5,0)
D=((1,-.5,-.5),(0,math.sqrt(3)/2,-math.sqrt(3)/2))
def vec(x):
 if not x or any(type(v) is bool or not isinstance(v,Real) or not math.isfinite(v) for v in x):raise ValueError('Nonempty finite real vector required')
 return list(x)
def dot(x,y):
 vec(x);vec(y)
 if len(x)!=len(y):raise ValueError('Incompatible lengths')
 return sum(a*b for a,b in zip(x,y))
def mv(matrix,x):
 if not matrix:raise ValueError('Empty matrix')
 return [dot(row,x) for row in matrix]
def add(x,y):
 dot(x,y);return [a+b for a,b in zip(x,y)]
def relu(x):return [max(0,v) for v in vec(x)]
def mlp(x,ablate=None):
 x=vec(x)
 if len(x)!=3 or ablate is not None and (type(ablate) is not int or not 0<=ablate<4):raise ValueError('Invalid input or unit index')
 z=add(mv(WIN,x),BIN);h=relu(z);natural=h[:]
 if ablate is not None:h[ablate]=0
 contributions=[[h[j]*row[j] for row in WOUT] for j in range(4)]
 delta=add(mv(WOUT,h),BOUT);y=add(x,delta)
 return {'x':x,'z':z,'natural_h':natural,'h':h,'column_contributions':contributions,'delta':delta,'y':y,'score':y[0]+y[2]}
def dictionary(a):
 a=vec(a)
 if len(a)!=3 or min(a)<0:raise ValueError('Three nonnegative strengths required')
 r=mv(D,a);pre=mv(list(zip(*D)),r);reconstruction=relu(pre)
 return {'strengths':a,'representation':r,'pre_relu':pre,'reconstruction':reconstruction,'squared_error':sum((x-y)**2 for x,y in zip(a,reconstruction))}
if __name__=='__main__':
 p=[1/math.sqrt(2)]*2;q=[1/math.sqrt(2),-1/math.sqrt(2)];r=add([2*x for x in p],q)
 print(json.dumps({'python':platform.python_version(),'lesson_mlp':mlp([1,-1,2]),'lab_mlp':mlp([2,0,1]),
  'ablations':{str(i):mlp([2,0,1],i) for i in [0,2,1]},'dictionary_cases':[dictionary(a) for a in [[0,2,0],[1,0,1],[1,2,0],[1,1,1]]],
  'orthogonal_code':{'representation':r,'recovered':[dot(p,r),dot(q,r)]}},indent=2))
