stack-attention-pytorch 0.0.2__tar.gz → 0.0.3__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.3}/PKG-INFO +32 -1
- {stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.3}/README.md +31 -0
- {stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.3}/pyproject.toml +1 -1
- {stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.3}/stack_attention/stack_trans_layer.py +19 -5
- {stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.3}/.gitignore +0 -0
- {stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.3}/LICENSE +0 -0
- {stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.3}/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.3
|
|
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,37 @@ 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(
|
|
67
|
+
tokens,
|
|
68
|
+
stochastic_action = stochastic_action,
|
|
69
|
+
hard_action = hard_action
|
|
70
|
+
)
|
|
71
|
+
|
|
72
|
+
out2, state = layer(
|
|
73
|
+
tokens,
|
|
74
|
+
stochastic_action = stochastic_action,
|
|
75
|
+
hard_action = hard_action
|
|
76
|
+
)
|
|
77
|
+
|
|
78
|
+
assert out1.shape == out2.shape == tokens.shape
|
|
79
|
+
```
|
|
80
|
+
|
|
50
81
|
## Citations
|
|
51
82
|
|
|
52
83
|
```bibtex
|
|
@@ -2,6 +2,37 @@
|
|
|
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(
|
|
22
|
+
tokens,
|
|
23
|
+
stochastic_action = stochastic_action,
|
|
24
|
+
hard_action = hard_action
|
|
25
|
+
)
|
|
26
|
+
|
|
27
|
+
out2, state = layer(
|
|
28
|
+
tokens,
|
|
29
|
+
stochastic_action = stochastic_action,
|
|
30
|
+
hard_action = hard_action
|
|
31
|
+
)
|
|
32
|
+
|
|
33
|
+
assert out1.shape == out2.shape == tokens.shape
|
|
34
|
+
```
|
|
35
|
+
|
|
5
36
|
## Citations
|
|
6
37
|
|
|
7
38
|
```bibtex
|
{stack_attention_pytorch-0.0.2 → stack_attention_pytorch-0.0.3}/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
|
+
learned_residual_gate &= add_residual
|
|
100
|
+
|
|
101
|
+
self.residual_scale = Parameter(tensor(0.)) if 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
|
|
|
@@ -188,7 +197,7 @@ class StackTransLayer(Module):
|
|
|
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
202
|
out = rearrange(stack_inputs, 'b h d -> b (h d)')
|
|
194
203
|
|
|
@@ -196,4 +205,9 @@ class StackTransLayer(Module):
|
|
|
196
205
|
|
|
197
206
|
out = inverse_pack(out)
|
|
198
207
|
|
|
208
|
+
# maybe add residual
|
|
209
|
+
|
|
210
|
+
if self.add_residual:
|
|
211
|
+
out = out + residual * self.residual_scale.exp()
|
|
212
|
+
|
|
199
213
|
return out, next_stack_states
|
|
File without changes
|
|
File without changes
|
|
File without changes
|