Transformer Simple Example

This example extends the example of building a neural network from scratch using Pytorch. In it, we have one input, and one output. (that is, to simplify we only try to train the network to recognize a single input)

The input is a single column vector
{% x = \begin{bmatrix} 1 \\ 1 \\ \end{bmatrix} %}
which should map to the output
{% Y = \begin{bmatrix} 5 \\ 2 \\ \end{bmatrix} %}


We follow the basic formula for creating a transformer.
{% softmax[ XW^q (W^k)^T X^T] X W^v %}

Q,K and V

In order to compute the intermediate values, {% Q,K,V %}, we construct the corresponding weight matrces.
Wq = [[1.0,1.0,1.0]] Wk = [[1.0,2.0,1.0]] Wv = [[1.0]] wq = torch.tensor(Wq, requires_grad=True) wk = torch.tensor(Wk, requires_grad=True) wv = torch.tensor(Wv, requires_grad=True)

Then, within the model computation, we have the following

Q = input @ wq K = input @ wk V = input @ wv

Full Script

import torch 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]] 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) 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 Y = torch.softmax(Q @ K.T, dim=0) @ V # 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))