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))