stack-attention-pytorch 0.0.3__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.3 → stack_attention_pytorch-0.0.4}/PKG-INFO +3 -8
- {stack_attention_pytorch-0.0.3 → stack_attention_pytorch-0.0.4}/README.md +2 -7
- {stack_attention_pytorch-0.0.3 → stack_attention_pytorch-0.0.4}/pyproject.toml +1 -1
- {stack_attention_pytorch-0.0.3 → stack_attention_pytorch-0.0.4}/stack_attention/stack_trans_layer.py +6 -5
- {stack_attention_pytorch-0.0.3 → stack_attention_pytorch-0.0.4}/.gitignore +0 -0
- {stack_attention_pytorch-0.0.3 → stack_attention_pytorch-0.0.4}/LICENSE +0 -0
- {stack_attention_pytorch-0.0.3 → 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
|
|
@@ -63,16 +63,11 @@ tokens = torch.randn(2, 512, 256)
|
|
|
63
63
|
|
|
64
64
|
layer = StackTransLayer(256)
|
|
65
65
|
|
|
66
|
-
out1, state = layer(
|
|
67
|
-
tokens,
|
|
68
|
-
stochastic_action = stochastic_action,
|
|
69
|
-
hard_action = hard_action
|
|
70
|
-
)
|
|
66
|
+
out1, state = layer(tokens)
|
|
71
67
|
|
|
72
68
|
out2, state = layer(
|
|
73
69
|
tokens,
|
|
74
|
-
|
|
75
|
-
hard_action = hard_action
|
|
70
|
+
stack_states = state
|
|
76
71
|
)
|
|
77
72
|
|
|
78
73
|
assert out1.shape == out2.shape == tokens.shape
|
|
@@ -18,16 +18,11 @@ tokens = torch.randn(2, 512, 256)
|
|
|
18
18
|
|
|
19
19
|
layer = StackTransLayer(256)
|
|
20
20
|
|
|
21
|
-
out1, state = layer(
|
|
22
|
-
tokens,
|
|
23
|
-
stochastic_action = stochastic_action,
|
|
24
|
-
hard_action = hard_action
|
|
25
|
-
)
|
|
21
|
+
out1, state = layer(tokens)
|
|
26
22
|
|
|
27
23
|
out2, state = layer(
|
|
28
24
|
tokens,
|
|
29
|
-
|
|
30
|
-
hard_action = hard_action
|
|
25
|
+
stack_states = state
|
|
31
26
|
)
|
|
32
27
|
|
|
33
28
|
assert out1.shape == out2.shape == tokens.shape
|
{stack_attention_pytorch-0.0.3 → stack_attention_pytorch-0.0.4}/stack_attention/stack_trans_layer.py
RENAMED
|
@@ -96,9 +96,9 @@ class StackTransLayer(Module):
|
|
|
96
96
|
# maybe combining with residual
|
|
97
97
|
|
|
98
98
|
self.add_residual = add_residual
|
|
99
|
-
learned_residual_gate
|
|
99
|
+
self.learned_residual_gate = learned_residual_gate and add_residual
|
|
100
100
|
|
|
101
|
-
self.residual_scale = Parameter(tensor(0.)) if learned_residual_gate else None
|
|
101
|
+
self.residual_scale = Parameter(tensor(0.)) if self.learned_residual_gate else None
|
|
102
102
|
|
|
103
103
|
def forward(
|
|
104
104
|
self,
|
|
@@ -191,7 +191,7 @@ class StackTransLayer(Module):
|
|
|
191
191
|
read_stack_attn_logits = self.to_read_stack_attn(next_stack_with_null)
|
|
192
192
|
read_stack_attn_logits = rearrange(read_stack_attn_logits, '... 1 -> ...')
|
|
193
193
|
|
|
194
|
-
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))
|
|
195
195
|
|
|
196
196
|
read_stack_attn = read_stack_attn_logits.softmax(dim = -1)
|
|
197
197
|
|
|
@@ -199,7 +199,7 @@ class StackTransLayer(Module):
|
|
|
199
199
|
|
|
200
200
|
# combine heads
|
|
201
201
|
|
|
202
|
-
out = rearrange(
|
|
202
|
+
out = rearrange(read_stack_out, 'b h d -> b (h d)')
|
|
203
203
|
|
|
204
204
|
out = self.combine(out)
|
|
205
205
|
|
|
@@ -208,6 +208,7 @@ class StackTransLayer(Module):
|
|
|
208
208
|
# maybe add residual
|
|
209
209
|
|
|
210
210
|
if self.add_residual:
|
|
211
|
-
|
|
211
|
+
residual_scale = self.residual_scale.exp() if exists(self.residual_scale) else 1.
|
|
212
|
+
out = out + residual * residual_scale
|
|
212
213
|
|
|
213
214
|
return out, next_stack_states
|
|
File without changes
|
|
File without changes
|
|
File without changes
|