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

lr = 0.01
adaptrate = 0.03

class DenseConnection:
	def __init__(self, w,scale):
		self.w=w
		self.scale = scale
	def make(in_size,out_size,scale=1.0):
		return DenseConnection(torch.rand(in_size,out_size),scale/in_size)
	def make2(in_size,out_size,scale=1.0):
		w=torch.rand(in_size,out_size)
		return DenseConnection(w,scale/in_size),DenseConnection(w.T,scale/out_size)
	def forward(self, x): return x@self.w*self.scale
	def updateWeight(self, sender,receiver):
		change = sender.act[:,None]*receiver.act[None,:] - sender.minusPhaseAct[:,None]*receiver.minusPhaseAct[None,:]
		self.w += change*lr
		self.w.clamp(0,1)
		self.w += (torch.where(receiver.targetDiff[None,:]<0.0, 0.0, 1.0) - self.w) * (receiver.targetDiff[None,:].abs()) * adaptrate
	def updateWeightSpecific(self, senderMinusPhase,senderPlusPhase,receiver):
		change = senderPlusPhase[:,None]*receiver.act[None,:] - senderMinusPhase[:,None]*receiver.minusPhaseAct[None,:]
		self.w += change*lr
		self.w.clamp(0,1)
		self.w += (torch.where(receiver.targetDiff[None,:]<0.0, 0.0, 1.0) - self.w) * (receiver.targetDiff[None,:].abs()) * adaptrate


class LocalPoolConnection:
	"""
	Supports different strides, different pool sizes, and different neuron counts 
	in sender and receiver layers.
	
	Structure:
	- Sender: (S_px, S_py, S_neurons)
	- Receiver: (R_px, R_py, R_neurons)
	- Strides: (stride_x, stride_y) in terms of receiver pools per sender pool.
	  Actually, it's easier to define: 
	  Sender pool (i, j) connects to Receiver pools in a local neighborhood.
	  But usually, local connection means Receiver pool (r, c) connects to 
	  Sender pools (r+dx, c+dy).
	  
	Let's define strides as:
	- Sender stride: distance between sender pool centers in terms of sender pool grid.
	- Receiver stride: distance between receiver pool centers.
	
	Actually, standard "local connection" with different strides:
	If stride_x=1, receiver pool r connects to sender pools r, r+1, r+2...
	If stride_x=2, receiver pool r connects to sender pools r*2, r*2+1, r*2+2...
	
	Let's use:
	- Sender grid size: (S_px, S_py)
	- Receiver grid size: (R_px, R_py)
	- Neurons per sender pool: S_n
	- Neurons per receiver pool: R_n
	- Neighborhood size: neighborhood_size (e.g., 1 for 3x3, 2 for 5x5)
	
	For each receiver pool (rx, ry), it connects to sender pools:
	(rx * stride_x + dx, ry * stride_y + dy) for dx, dy in neighborhood
	
	Where stride_x, stride_y are the strides.
	"""
	def __init__(self, w, scale, sender_indices, receiver_indices, 
				 s_px, s_py, s_neurons, r_px, r_py, r_neurons, 
				 stride_x, stride_y, neighborhood_size=1):
		"""
		Args:
			w: 1D tensor of weights
			scale: scaling factor
			sender_indices: 1D long tensor, sender neuron indices
			receiver_indices: 1D long tensor, receiver neuron indices
			s_px, s_py: sender grid dimensions
			s_neurons: neurons per sender pool
			r_px, r_py: receiver grid dimensions
			r_neurons: neurons per receiver pool
			stride_x, stride_y: strides
			neighborhood_size: size of neighborhood
		"""
		self.w = w
		self.scale = scale
		self.sender_indices = sender_indices
		self.receiver_indices = receiver_indices
		self.s_px = s_px
		self.s_py = s_py
		self.s_neurons = s_neurons
		self.r_px = r_px
		self.r_py = r_py
		self.r_neurons = r_neurons
		self.stride_x = stride_x
		self.stride_y = stride_y
		self.neighborhood_size = neighborhood_size
		self.neighborhood_offsets = [(dx, dy) for dx in [-neighborhood_size, neighborhood_size+1] for dy in [-neighborhood_size, neighborhood_size+1]]
		# Fix: neighborhood_offsets should be correct
		self.neighborhood_offsets = []
		for dx in range(-neighborhood_size, neighborhood_size+1):
			for dy in range(-neighborhood_size, neighborhood_size+1):
				self.neighborhood_offsets.append((dx, dy))

	@classmethod
	def make(cls, s_px, s_py, s_neurons, r_px, r_py, r_neurons, 
			 stride_x=1, stride_y=1, scale=1.0, neighborhood_size=1):
		"""
		Creates a local pool connection.
		"""
		offsets = [(dx, dy) for dx in range(-neighborhood_size, neighborhood_size+1) 
				   for dy in range(-neighborhood_size, neighborhood_size+1)]
		
		sender_indices_list = []
		receiver_indices_list = []
		
		for rx in range(r_px):
			for ry in range(r_py):
				for dx, dy in offsets:
					# Sender pool coordinates
					spx = rx * stride_x + dx
					spy = ry * stride_y + dy
					
					if 0 <= spx < s_px and 0 <= spy < s_py:
						# Connect each neuron in sender pool to each neuron in receiver pool
						# This creates a dense local connection within the neighborhood
						# But if s_neurons != r_neurons, we need to handle it.
						# Let's assume full connectivity between sender and receiver neurons in the neighborhood
						# Or pointwise? The original code did pointwise: n sender -> n receiver
						# So we connect sender neuron (spx, spy, n) to receiver neuron (rx, ry, n)
						
						for n in range(s_neurons):
							# Calculate base indices for the pool connection
							base_sender_idx = spx * (s_py * s_neurons) + spy * s_neurons
							base_receiver_idx = rx * (r_py * r_neurons) + ry * r_neurons
							
							# Add indices for all receiver neurons for each sender neuron
							for r_n in range(r_neurons):
								sender_indices_list.append(base_sender_idx + n)
								receiver_indices_list.append(base_receiver_idx + r_n)
		
		sender_indices = torch.tensor(sender_indices_list, dtype=torch.long)
		receiver_indices = torch.tensor(receiver_indices_list, dtype=torch.long)
		
		w = torch.rand(len(sender_indices_list))
		
		# Calculate scale factor
		total_receiver_neurons = r_px * r_py * r_neurons
		fan_in = torch.zeros(total_receiver_neurons, dtype=torch.float)
		fan_in.index_add_(0, receiver_indices, torch.ones_like(w))
		avg_fan_in = fan_in.mean()
		
		scale_factor = scale / avg_fan_in if avg_fan_in > 0 else scale
		
		return cls(w, scale_factor, sender_indices, receiver_indices, 
				   s_px, s_py, s_neurons, r_px, r_py, r_neurons, 
				   stride_x, stride_y, neighborhood_size)

	@classmethod
	def make2(cls, s_px, s_py, s_neurons, r_px, r_py, r_neurons, 
			 stride_x=1, stride_y=1, scale=1.0, neighborhood_size=1):
		"""
		Creates transposed connection.
		"""
		conn = cls.make(s_px, s_py, s_neurons, r_px, r_py, r_neurons, stride_x, stride_y, scale, neighborhood_size)
		# Transposed: swap sender and receiver
		# But we need to adjust strides and grid sizes
		# If original: receiver pool rx connects to sender pool rx*stride + dx
		# Transposed: sender pool (rx*stride + dx) connects to receiver pool rx
		
		# New receiver (old sender) grid: (s_px, s_py)
		# New sender (old receiver) grid: (r_px, r_py)
		# New strides: inverse of original strides?
		
		# Actually, let's just swap indices and keep the structure.
		# The connectivity graph is reversed.
		
		new_sender_indices = conn.receiver_indices
		new_receiver_indices = conn.sender_indices
		new_w = conn.w.clone()
		
		# New grid dimensions
		new_r_px, new_r_py = conn.s_px, conn.s_py
		new_s_px, new_s_py = conn.r_px, conn.r_py
		new_r_neurons = conn.s_neurons
		new_s_neurons = conn.r_neurons
		
		# New strides: 
		# If original stride was sx, sy, the transposed stride is 1/sx, 1/sy?
		# Not exactly, because the grid sizes are different.
		# Let's keep the same neighborhood size and offsets.
		
		# Recalculate scale
		total_new_receiver_neurons = new_r_px * new_r_py * new_r_neurons
		fan_in = torch.zeros(total_new_receiver_neurons, dtype=torch.float)
		fan_in.index_add_(0, new_receiver_indices, torch.ones_like(new_w))
		avg_fan_in = fan_in.mean()
		
		scale_factor = conn.scale / (conn.scale / (conn.w.shape[0] / total_new_receiver_neurons if total_new_receiver_neurons > 0 else 1)) 
		# Simpler: use the same scale factor logic
		scale_factor = conn.scale / avg_fan_in if avg_fan_in > 0 else conn.scale
		
		new_offsets = conn.neighborhood_offsets
		
		return conn,cls(new_w, scale_factor, new_sender_indices, new_receiver_indices,
				   new_s_px, new_s_py, new_s_neurons, new_r_px, new_r_py, new_r_neurons,
				   conn.stride_x, conn.stride_y, conn.neighborhood_size)

	def forward(self, x):
		"""
		Forward pass.
		x: sender activations, shape (S_px * S_py * S_n,)
		Returns: receiver activations, shape (R_px * R_py * R_n,)
		"""
		assert len(x.shape)==1 and x.shape[0]==self.s_px*self.s_py*self.s_neurons, x.shape
		sender_values = x[self.sender_indices]
		weighted_values = self.w * sender_values
		receiver_flat = torch.zeros(self.r_px * self.r_py * self.r_neurons)
		receiver_flat.index_add_(0, self.receiver_indices, weighted_values)
		return receiver_flat * self.scale

	def updateWeight(self, sender, receiver):
		"""
		Update weights.
		sender: object with .act, .minusPhaseAct
		receiver: object with .act, .minusPhaseAct, .targetDiff
		"""
		# Flatten activations
		if sender.act.dim() > 1:
			sender_plus = sender.act.view(sender.act.shape[0], -1)
			sender_minus = sender.minusPhaseAct.view(sender.minusPhaseAct.shape[0], -1)
		else:
			sender_plus = sender.act.unsqueeze(0)
			sender_minus = sender.minusPhaseAct.unsqueeze(0)
		
		if receiver.act.dim() > 1:
			receiver_act = receiver.act.view(receiver.act.shape[0], -1)
			receiver_minus = receiver.minusPhaseAct.view(receiver.minusPhaseAct.shape[0], -1)
			target_diff = receiver.targetDiff.view(receiver.targetDiff.shape[0], -1)
		else:
			receiver_act = receiver.act.unsqueeze(0)
			receiver_minus = receiver.minusPhaseAct.unsqueeze(0)
			target_diff = receiver.targetDiff.unsqueeze(0)
		
		sender_plus_vals = sender_plus[:, self.sender_indices]
		receiver_act_vals = receiver_act[:, self.receiver_indices]
		sender_minus_vals = sender_minus[:, self.sender_indices]
		receiver_minus_vals = receiver_minus[:, self.receiver_indices]
		
		change_plus = sender_plus_vals * receiver_act_vals
		change_minus = sender_minus_vals * receiver_minus_vals
		change = (change_plus - change_minus).mean(dim=0)
		
		self.w += change * lr
		self.w.clamp_(0, 1)
		
		target_diff_vals = target_diff[:, self.receiver_indices]
		target_mask = torch.where(target_diff_vals < 0, 0.0, 1.0)
		target_update = (target_mask - self.w) * target_diff_vals.abs() * adaptrate
		self.w += target_update.flatten()

	def updateWeightSpecific(self, senderMinusPhase, senderPlusPhase, receiver):
		# Flatten activations
		if senderPlusPhase.dim() > 1:
			sender_plus = senderPlusPhase.view(senderPlusPhase.shape[0], -1)
			sender_minus = senderMinusPhase.view(senderMinusPhase.shape[0], -1)
		else:
			sender_plus = senderPlusPhase.unsqueeze(0)
			sender_minus = senderMinusPhase.unsqueeze(0)
		
		if receiver.act.dim() > 1:
			receiver_act = receiver.act.view(receiver.act.shape[0], -1)
			receiver_minus = receiver.minusPhaseAct.view(receiver.minusPhaseAct.shape[0], -1)
			target_diff = receiver.targetDiff.view(receiver.targetDiff.shape[0], -1)
		else:
			receiver_act = receiver.act.unsqueeze(0)
			receiver_minus = receiver.minusPhaseAct.unsqueeze(0)
			target_diff = receiver.targetDiff.unsqueeze(0)
		
		sender_plus_vals = sender_plus[:, self.sender_indices]
		receiver_act_vals = receiver_act[:, self.receiver_indices]
		sender_minus_vals = sender_minus[:, self.sender_indices]
		receiver_minus_vals = receiver_minus[:, self.receiver_indices]
		
		change_plus = sender_plus_vals * receiver_act_vals
		change_minus = sender_minus_vals * receiver_minus_vals
		change = (change_plus - change_minus).mean(dim=0)
		
		self.w += change * lr
		self.w.clamp_(0, 1)
		
		target_diff_vals = target_diff[:, self.receiver_indices]
		target_mask = torch.where(target_diff_vals < 0, 0.0, 1.0)
		target_update = (target_mask - self.w) * target_diff_vals.abs() * adaptrate
		self.w += target_update.flatten()

class Layer:
	def __init__(self, shape):
		self.shape = shape
		if isinstance(shape, tuple):
			prod = 1
			for s in shape: prod *= s
			self.size = prod
		else:
			self.size = shape
			self.shape = (shape,)
		self.minusPhaseAct=None
		self.act=None
		self.prevAct=torch.zeros(self.size)
		self.avgAct = torch.zeros(self.size)
	def update(self, input):
		feedforwardInhibition = input.mean()*0.7 + input.amax()*0.3
		if len(self.shape)>1:
			inPoolDim=len(self.shape)-1
			inputView = input.reshape(self.shape)
			feedforwardInhibitionPool = inputView.mean(inPoolDim,keepdim=True)*0.7 + inputView.amax(inPoolDim,keepdim=True)*0.3
			feedforwardInhibition = torch.maximum(feedforwardInhibition,feedforwardInhibitionPool.expand(self.shape).flatten())
		#self.feedforwardInhibition=feedforwardInhibition
		input = (input - feedforwardInhibition)
		self.act = (input*64.0).tanh_().clamp_(0)
		return self.act
	def updatePlusPhase(self, input):
		self.update(input)
		#input = (input - self.feedforwardInhibition)
		#self.act = (input*64.0).tanh_().clamp_(0)
		return self.act
	def slowUpdate(self):
		self.avgAct += (self.act-self.avgAct)*adaptrate
		inPoolDim=len(self.shape)-1
		self.targetDiff = self.avgAct.reshape(self.shape).mean(inPoolDim,keepdim=True).expand(self.shape).flatten() - self.avgAct

inpLayer=Layer((5,5,3))
V1=Layer((5,5,16))
V2=Layer((5,5,9))
V3=Layer((5,5,4))

inpp=Layer(inpLayer.shape)
V1p=Layer(V1.shape)
V2p=Layer(V2.shape)

inpToV1,V1ToInp=LocalPoolConnection.make2(5,5,3, 5,5,16)
v1ToV2,V2ToV1=LocalPoolConnection.make2(5,5,16, 5,5,9)
v2ToV3,V3ToV2=LocalPoolConnection.make2(5,5,9, 5,5,4)

v3ToV2p=LocalPoolConnection.make(5,5,4, 5,5,9)
v2ToV1p=LocalPoolConnection.make(5,5,9, 5,5,16)
v1ToInpp=LocalPoolConnection.make(5,5,16, 5,5,3)

inppToV1=LocalPoolConnection.make(5,5,3, 5,5,16, scale=0.1)
v1pToV2=LocalPoolConnection.make(5,5,16, 5,5,9, scale=0.1)
v2pToV3=LocalPoolConnection.make(5,5,9, 5,5,4, scale=0.1)

i=0

fig, axs = plt.subplots(4,3)
plt.tight_layout(pad=0)
plt.subplots_adjust(wspace=0.05, hspace=0.05)
for ax in axs.flat: ax.tick_params(labelbottom=False, labelleft=False)
def showAt(ax, v):
	ax.clear()
	ax.axis('off')
	if v is None: return
	if len(v.shape)==1: v=v.unsqueeze(0)
	if len(v.shape)==3: v=v.flatten(start_dim=1)
	ax.imshow(v, vmin=0, vmax=1)


def draw_frame(time_step, grid_dim=100):
	# Calculate position using sine for smooth oscillation
	x = int((grid_dim / 2) * (1 + torch.sin(time_step * 0.1)))
	y = int((grid_dim / 2) * (1 + torch.cos(time_step * 0.07)))
	
	# Create empty frame
	frame = torch.zeros(grid_dim, grid_dim, 3)
	
	# Draw single pixel
	if 0 <= x < grid_dim and 0 <= y < grid_dim:
		frame[y, x, :] = 1.0
		
	return frame

  
  
  
for j in range(1000):
	inp=draw_frame(torch.tensor(i*3.0),5).flatten()
	i+=1

	inppa=inpp.update(v1ToInpp.forward(V1.prevAct))
	v1pa=V1p.update(v2ToV1p.forward(V2.prevAct))
	v2pa=V2p.update(v3ToV2p.forward(V3.prevAct))

	inpa=inpLayer.act = inp
	v1a=V1.update(inpToV1.forward(inpa)+inppToV1.forward(inppa))
	v2a=V2.update(v1ToV2.forward(v1a)+v1pToV2.forward(v1pa))
	v3a=V3.update(v2ToV3.forward(v2a)+v2pToV3.forward(v2pa))
	#spread backward so it spreads more completely
	v2a=V2.update(V3ToV2.forward(v3a)+v1ToV2.forward(v1a)+v1pToV2.forward(v1pa))
	v1a=V1.update(V2ToV1.forward(v2a)+inpToV1.forward(inpa)+inppToV1.forward(inppa))

	# do again but clamp p layers

	inpLayer.minusPhaseAct = inpLayer.act
	V1.minusPhaseAct = V1.act
	V2.minusPhaseAct = V2.act
	V3.minusPhaseAct = V3.act

	inpp.minusPhaseAct=inpp.act
	V1p.minusPhaseAct=V1p.act
	V2p.minusPhaseAct=V2p.act

	inppaPlusPhase=inpp.act=inpa
	v1paPlusPhase=V1p.act=v1a
	v2paPlusPhase=V2p.act=v2a

	inpa = inpLayer.act = inp
	v1aPlusPhase=V1.updatePlusPhase(inpToV1.forward(inpa)+inppToV1.forward(inppaPlusPhase))
	v2aPlusPhase=V2.updatePlusPhase(v1ToV2.forward(v1aPlusPhase)+v1pToV2.forward(v1paPlusPhase))
	v3aPlusPhase=V3.updatePlusPhase(v2ToV3.forward(v2aPlusPhase)+v2pToV3.forward(v2paPlusPhase))

	v2aPlusPhase=V2.updatePlusPhase(V3ToV2.forward(v3aPlusPhase)+v1ToV2.forward(v1aPlusPhase)+v1pToV2.forward(v1paPlusPhase))
	v1aPlusPhase=V1.updatePlusPhase(V2ToV1.forward(v2aPlusPhase)+inpToV1.forward(inpa)+inppToV1.forward(inppaPlusPhase))

	#update averages
	inpLayer.slowUpdate()
	V1.slowUpdate()
	V2.slowUpdate()
	V3.slowUpdate()

	inpp.slowUpdate()
	V1p.slowUpdate()
	V2p.slowUpdate()

	#update weight

	inpToV1.updateWeight(inpLayer,V1)
	v1ToV2.updateWeight(V1,V2)
	V2ToV1.updateWeight(V2,V1)
	v2ToV3.updateWeight(V2,V3)
	V3ToV2.updateWeight(V3,V2)

	v1ToInpp.updateWeightSpecific(V1.prevAct,V1.prevAct,inpp)
	v2ToV1p.updateWeightSpecific(V2.prevAct,V2.prevAct,V1p)
	v3ToV2p.updateWeightSpecific(V3.prevAct,V3.prevAct,V2p)

	inppToV1.updateWeight(inpp,V1)
	v1pToV2.updateWeight(V1p,V2)
	v2pToV3.updateWeight(V2p,V3)

	#update prevAct
	inpLayer.prevAct = inpLayer.act
	V1.prevAct = V1.act
	V2.prevAct = V2.act
	V3.prevAct = V3.act

	inpp.prevAct = inpp.act
	V1p.prevAct = V1p.act
	V2p.prevAct = V2p.act

	showAt(axs[0,0], inp.view(5,5,3))
	showAt(axs[0,1], inppa.view(5,5,3))
	showAt(axs[1,0], v1a.view(5,5,16))
	showAt(axs[1,1], v1pa.view(5,5,16))
	showAt(axs[1,2], (V1.act-V1.minusPhaseAct).view(5,5,16))
	showAt(axs[2,0], v2a.view(5,5,9))
	showAt(axs[2,1], v2pa.view(5,5,9))
	showAt(axs[3,0], v3a.view(5,5,4))
	plt.pause(1)
plt.pause(100000)