stack-attention-pytorch 0.0.2__tar.gz → 0.0.4__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: stack-attention-pytorch
3
- Version: 0.0.2
3
+ Version: 0.0.4
4
4
  Summary: Implementation of Differentiable Stacks for augmenting Transformers / Attention
5
5
  Project-URL: Homepage, https://pypi.org/project/stack-attention-pytorch/
6
6
  Project-URL: Repository, https://github.com/lucidrains/stack-attention-pytorch
@@ -47,6 +47,32 @@ Description-Content-Type: text/markdown
47
47
 
48
48
  For following a line of research that augments attention with a differentiable stack, beginning with DuSell et al. at ETH Zurich
49
49
 
50
+ ## Install
51
+
52
+ ```bash
53
+ $ pip install stack-attention
54
+ ```
55
+
56
+ ## Usage
57
+
58
+ ```python
59
+ import torch
60
+ from stack_attention.stack_trans_layer import StackTransLayer
61
+
62
+ tokens = torch.randn(2, 512, 256)
63
+
64
+ layer = StackTransLayer(256)
65
+
66
+ out1, state = layer(tokens)
67
+
68
+ out2, state = layer(
69
+ tokens,
70
+ stack_states = state
71
+ )
72
+
73
+ assert out1.shape == out2.shape == tokens.shape
74
+ ```
75
+
50
76
  ## Citations
51
77
 
52
78
  ```bibtex
@@ -2,6 +2,32 @@
2
2
 
3
3
  For following a line of research that augments attention with a differentiable stack, beginning with DuSell et al. at ETH Zurich
4
4
 
5
+ ## Install
6
+
7
+ ```bash
8
+ $ pip install stack-attention
9
+ ```
10
+
11
+ ## Usage
12
+
13
+ ```python
14
+ import torch
15
+ from stack_attention.stack_trans_layer import StackTransLayer
16
+
17
+ tokens = torch.randn(2, 512, 256)
18
+
19
+ layer = StackTransLayer(256)
20
+
21
+ out1, state = layer(tokens)
22
+
23
+ out2, state = layer(
24
+ tokens,
25
+ stack_states = state
26
+ )
27
+
28
+ assert out1.shape == out2.shape == tokens.shape
29
+ ```
30
+
5
31
  ## Citations
6
32
 
7
33
  ```bibtex
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "stack-attention-pytorch"
3
- version = "0.0.2"
3
+ version = "0.0.4"
4
4
  description = "Implementation of Differentiable Stacks for augmenting Transformers / Attention"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -2,7 +2,7 @@ from __future__ import annotations
2
2
  from functools import partial
3
3
 
4
4
  import torch
5
- from torch import cat, nn, Tensor
5
+ from torch import cat, nn, Tensor, tensor
6
6
  from torch.nn import Module, Linear, RMSNorm, Parameter
7
7
  import torch.nn.functional as F
8
8
 
@@ -58,7 +58,9 @@ class StackTransLayer(Module):
58
58
  num_stacks = 4, # heads
59
59
  dim_stack = 16, # dim head
60
60
  stack_size = 24,
61
- prenorm = False
61
+ prenorm = False,
62
+ add_residual = True,
63
+ learned_residual_gate = True
62
64
  ):
63
65
  super().__init__()
64
66
 
@@ -87,10 +89,17 @@ class StackTransLayer(Module):
87
89
 
88
90
  self.null_stack = Parameter(torch.randn(dim_stack) * 1e-2)
89
91
 
90
- # combine
92
+ # combining stack reads across number of stacks
91
93
 
92
94
  self.combine = LinearNoBias(dim_inner, dim)
93
95
 
96
+ # maybe combining with residual
97
+
98
+ self.add_residual = add_residual
99
+ self.learned_residual_gate = learned_residual_gate and add_residual
100
+
101
+ self.residual_scale = Parameter(tensor(0.)) if self.learned_residual_gate else None
102
+
94
103
  def forward(
95
104
  self,
96
105
  tokens,
@@ -101,7 +110,7 @@ class StackTransLayer(Module):
101
110
  ):
102
111
  assert action_temperature > 0.
103
112
 
104
- orig, device = tokens, tokens.device
113
+ residual, device = tokens, tokens.device
105
114
 
106
115
  # maybe pre norm
107
116
 
@@ -182,18 +191,24 @@ class StackTransLayer(Module):
182
191
  read_stack_attn_logits = self.to_read_stack_attn(next_stack_with_null)
183
192
  read_stack_attn_logits = rearrange(read_stack_attn_logits, '... 1 -> ...')
184
193
 
185
- read_stack_attn_logits = read_stack_attn_logits.masked_fill(next_mask_with_null, mask_value(next_stack_with_null))
194
+ read_stack_attn_logits = read_stack_attn_logits.masked_fill(~next_mask_with_null, mask_value(next_stack_with_null))
186
195
 
187
196
  read_stack_attn = read_stack_attn_logits.softmax(dim = -1)
188
197
 
189
198
  read_stack_out = einsum(next_stack_with_null, read_stack_attn, 'b h s d, b h s -> b h d')
190
199
 
191
- # combine
200
+ # combine heads
192
201
 
193
- out = rearrange(stack_inputs, 'b h d -> b (h d)')
202
+ out = rearrange(read_stack_out, 'b h d -> b (h d)')
194
203
 
195
204
  out = self.combine(out)
196
205
 
197
206
  out = inverse_pack(out)
198
207
 
208
+ # maybe add residual
209
+
210
+ if self.add_residual:
211
+ residual_scale = self.residual_scale.exp() if exists(self.residual_scale) else 1.
212
+ out = out + residual * residual_scale
213
+
199
214
  return out, next_stack_states