DebasishDhal99 commited on
Commit
8038560
·
1 Parent(s): db4dc58

feat: added bare minimum rel-position embedding

Browse files
relative_pos_embedding/relative_pos_embedding.py ADDED
@@ -0,0 +1,125 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # In simple terms, we are calculating this
2
+ ### Sm,n = q_m^T * k_n + b_m-n ## here m and n are the positions of the query and key vectors respectively. The similarity score must depend on the relative position of the query and key vectors as well. b is some learned function of the relative position.
3
+
4
+ # Some minor lacuna pending here in the implementation, needs to be cleared.
5
+
6
+ import numpy as np
7
+
8
+ import numpy as np
9
+
10
+
11
+ def softmax(x, axis=-1):
12
+ x = x - np.max(x, axis=axis, keepdims=True)
13
+ exp_x = np.exp(x)
14
+ return exp_x / np.sum(exp_x, axis=axis, keepdims=True)
15
+
16
+
17
+ def shaw_relative_attention(
18
+ X,
19
+ Wq,
20
+ Wk,
21
+ Wv,
22
+ relative_key_embeddings,
23
+ relative_value_embeddings,
24
+ max_relative_position
25
+ ):
26
+ n_tokens, d_model = X.shape
27
+
28
+ # --------------------------------------------------
29
+ # 1. Relative positions
30
+ # --------------------------------------------------
31
+
32
+ positions = np.arange(n_tokens)
33
+
34
+ relative_positions = (
35
+ positions[:, None]
36
+ - positions[None, :]
37
+ )
38
+
39
+ relative_positions = np.clip(
40
+ relative_positions,
41
+ -max_relative_position,
42
+ max_relative_position
43
+ )
44
+
45
+ # Convert [-max, ..., +max] → [0, ..., 2*max]
46
+ relative_indices = (
47
+ relative_positions
48
+ + max_relative_position
49
+ )
50
+
51
+ # --------------------------------------------------
52
+ # 2. Query
53
+ # --------------------------------------------------
54
+
55
+ Q = X @ Wq
56
+
57
+ # --------------------------------------------------
58
+ # 3. Relative Key embeddings
59
+ # --------------------------------------------------
60
+
61
+ relative_key = (
62
+ relative_key_embeddings[
63
+ relative_indices
64
+ ]
65
+ )
66
+
67
+ # Shape:
68
+ # (n_tokens, n_tokens, d_model)
69
+
70
+ # x_n + relative positional embedding
71
+ K_input = (
72
+ X[None, :, :]
73
+ + relative_key
74
+ )
75
+
76
+ # Apply Wk
77
+ K_relative = K_input @ Wk
78
+
79
+ # --------------------------------------------------
80
+ # 4. Attention scores
81
+ # --------------------------------------------------
82
+
83
+ scores = np.einsum(
84
+ "md,mnd->mn",
85
+ Q,
86
+ K_relative
87
+ )
88
+
89
+ # --------------------------------------------------
90
+ # 5. Relative Value embeddings
91
+ # --------------------------------------------------
92
+
93
+ relative_value = (
94
+ relative_value_embeddings[
95
+ relative_indices
96
+ ]
97
+ )
98
+
99
+ V_input = (
100
+ X[None, :, :]
101
+ + relative_value
102
+ )
103
+
104
+ V_relative = V_input @ Wv
105
+
106
+ # --------------------------------------------------
107
+ # 6. Attention weights
108
+ # --------------------------------------------------
109
+
110
+ attention_weights = softmax(
111
+ scores,
112
+ axis=-1
113
+ )
114
+
115
+ # --------------------------------------------------
116
+ # 7. Weighted Values
117
+ # --------------------------------------------------
118
+
119
+ output = np.einsum(
120
+ "mn,mnd->md",
121
+ attention_weights,
122
+ V_relative
123
+ )
124
+
125
+ return output