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.
- {stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.4}/PKG-INFO +27 -1
- {stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.4}/README.md +26 -0
- {stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.4}/pyproject.toml +1 -1
- {stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.4}/stack_attention/stack_trans_layer.py +22 -7
- {stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.4}/.gitignore +0 -0
- {stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.4}/LICENSE +0 -0
- {stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.4}/stack_attention/__init__.py +0 -0
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
Metadata-Version: 2.5
|
|
2
2
|
Name: stack-attention-pytorch
|
|
3
|
-
Version: 0.0.
|
|
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
|
{stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.4}/stack_attention/stack_trans_layer.py
RENAMED
|
@@ -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
|
-
#
|
|
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
|
-
|
|
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(
|
|
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|