import torch
import math
x = [[1.0],[1.0]]
y = [[5.0],[2.0]]
W = [[1.0,1.0],[1.0,2.0]]
W2 = [[1.0,1.0],[1.0,2.0]]
k = 1
p=3
Wq = [[1.0,1.0,1.0]]
Wk = [[1.0,2.0,1.0]]
Wv = [[1.0]]
Wq2 = [[1.1,1.1,1.1]]
Wk2 = [[1.1,2.1,1.1]]
Wv2 = [[1.1]]
Wm = [[1.0], [1.0]]
input = torch.tensor(x)
weights = torch.tensor(W, requires_grad=True)
weights2 = torch.tensor(W2, requires_grad=True)
wq = torch.tensor(Wq, requires_grad=True)
wk = torch.tensor(Wk, requires_grad=True)
wv = torch.tensor(Wv, requires_grad=True)
wq2 = torch.tensor(Wq2, requires_grad=True)
wk2 = torch.tensor(Wk2, requires_grad=True)
wv2 = torch.tensor(Wv2, requires_grad=True)
wm = torch.tensor(Wm, requires_grad=True)
output = torch.tensor(y)
learning_rate = 0.01
optimizer = torch.optim.SGD([weights, weights2], lr=learning_rate)
# 3. Multi-step optimization loop
epochs = 1000
for step in range(epochs):
# Clear out old gradients from the previous step
optimizer.zero_grad()
Q = input @ wq
K = input @ wk
V = input @ wv
Y1 = torch.softmax(torch.mul(Q @ K.T, math.sqrt(k)), dim=0) @ V
Q2 = input @ wq2
K2 = input @ wk2
V2 = input @ wv2
Y2 = torch.softmax(torch.mul(Q2 @ K2.T, math.sqrt(k)), dim=0) @ V2
Y = torch.cat((Y1, Y2), dim=1) @ wm
# Forward pass: calculate prediction and loss
result = weights @ Y
layer1_output = torch.sigmoid(result)
layer2_output = weights2@layer1_output
error1 = output - layer2_output
loss = error1.T @ error1
#get total loss as a float
numeric_loss = loss.item()
# Backward pass: compute the gradients
loss.backward()
# Optimization step: update the weights using SGD math
optimizer.step()
print('final loss = '+str(numeric_loss))
print('final weights is '+str(weights))
print('final output is '+str(layer2_output))