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.
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: stack-attention-pytorch
3
- Version: 0.0.3
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
- stochastic_action = stochastic_action,
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
- stochastic_action = stochastic_action,
30
- hard_action = hard_action
25
+ stack_states = state
31
26
  )
32
27
 
33
28
  assert out1.shape == out2.shape == tokens.shape
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "stack-attention-pytorch"
3
- version = "0.0.3"
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" }
@@ -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 &= add_residual
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(stack_inputs, 'b h d -> b (h d)')
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
- out = out + residual * self.residual_scale.exp()
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