my code (not finished)

https://html.cafe/xbe3fcc53?k=fb08ba4d3662c0ddf5b4f40ecd759de18883f04d

import torch
from torch import tensor
import torch.nn as nn
import torch.nn.functional as F
import matplotlib.pyplot as plt

class con:
	def __init__(self, in_size,out_size):
		self.w=torch.rand(in_size,out_size)
	def forward(self, x): return [email protected]
	def backward(self, x): return [email protected]
	def multForWeights(self,in_x,out_x): return in_x[:,None]*out_x[None,:]

inpToV1=con(20,50)
v1ToV2=con(50,25)
v2ToV3=con(25,15)

v3ToV2p=con(15,25)
v2ToV1p=con(25,50)
v1ToInpp=con(25,50)

inppToV1=con(20,50)
v1pToV2=con(50,25)
v2pToV3=con(25,15)

prevV1a=torch.zeros(50)
prevV2a=torch.zeros(25)
prevV3a=torch.zeros(15)

for i in range(50):
	inp=torch.zeros(20)
	inp[i]=1.0

	inppa=v1ToInpp.forward(prevV1a)
	v1pa=v2ToV1p.forward(prevV2a)
	v2pa=v3ToV2p.forward(prevV3a)

	v1a=inpToV1.forward(inp)+inppToV1.forward(inppa)
	v2a=v1ToV2.forward(v1a)+v1pToV2.forward(v1pa)
	v3a=v2ToV3.forward(v2a)+v2pToV3.forward(v2pa)
	#spread backward so it spreads more completely
	v2a=v2ToV3.backward(v3a)+v1ToV2.forward(v1a)+v1pToV2.forward(v1pa)
	v1a=v1ToV2.backward(v2a)+inpToV1.forward(inp)+inppToV1.forward(inppa)

	# do again but clamp p layers

	inppaPlusPhase=inp
	v1paPlusPhase=v1a
	v2paPlusPhase=v2a

	v1aPlusPhase=inpToV1.forward(inp)+inppToV1.forward(inppaPlusPhase)
	v2aPlusPhase=v1ToV2.forward(v1aPlusPhase)+v1pToV2.forward(v1paPlusPhase)
	v3aPlusPhase=v2ToV3.forward(v2aPlusPhase)+v2pToV3.forward(v2paPlusPhase)

	v2aPlusPhase=v2ToV3.backward(v3aPlusPhase)+v1ToV2.forward(v1aPlusPhase)+v1pToV2.forward(v1paPlusPhase)
	v1aPlusPhase=v1ToV2.backward(v2aPlusPhase)+inpToV1.forward(inp)+inppToV1.forward(inppaPlusPhase)

	#update weight

	inpToV1.w+=inpToV1.multForWeights(inp,v1aPlusPhase)-inpToV1.multForWeights(inp,v1a)*0.01; inpToV1.w.clamp(0,1)
	v1ToV2.w+=v1ToV2.multForWeights(v1aPlusPhase,v2aPlusPhase)-v1ToV2.multForWeights(v1a,v2a)*0.01; v1ToV2.w.clamp(0,1)
	v2ToV3.w+=v2ToV3.multForWeights(v2aPlusPhase,v3aPlusPhase)-v2ToV3.multForWeights(v2a,v3a)*0.01; v2ToV3.w.clamp(0,1)

	v1ToInpp.w+=v1ToInpp.multForWeights(v1aPlusPhase,inppaPlusPhase)-v1ToInpp.multForWeights(v1a,inppa)*0.01; v1ToInpp.w.clamp(0,1)
	v2ToV1p.w+=v2ToV1p.multForWeights(v2aPlusPhase,v1paPlusPhase)-v2ToV1p.multForWeights(v2a,v1pa)*0.01; v2ToV1p.w.clamp(0,1)
	v3ToV2p.w+=v3ToV2p.multForWeights(v3aPlusPhase,v2paPlusPhase)-v3ToV2p.multForWeights(v3a,v2pa)*0.01; v3ToV2p.w.clamp(0,1)

	inppToV1.w+=inppToV1.multForWeights(inppa,v1aPlusPhase)-inppToV1.multForWeights(inppa,v1a)*0.01; inppToV1.w.clamp(0,1)
	v1pToV2.w+=v1pToV2.multForWeights(v1paPlusPhase,v2aPlusPhase)-v1pToV2.multForWeights(v1pa,v2a)*0.01; v1pToV2.w.clamp(0,1)
	v2pToV3.w+=v2pToV3.multForWeights(v2paPlusPhase,v3aPlusPhase)-v2pToV3.multForWeights(v2pa,v3a)*0.01; v2pToV3.w.clamp(0,1)

	prevV1a=v1aPlusPhase
	prevV2a=v2aPlusPhase
	prevV3a=v3aPlusPhase