-
Notifications
You must be signed in to change notification settings - Fork 3
Expand file tree
/
Copy pathwriting.py
More file actions
123 lines (86 loc) · 4.42 KB
/
Copy pathwriting.py
File metadata and controls
123 lines (86 loc) · 4.42 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
"""Writing layer"""
import math
from typing import Optional, Tuple, Callable, List
import torch
import torch.nn.functional
from functions.utility_functions import exp_convolve
from models.neuron_models import NeuronModel
class WritingLayer(torch.nn.Module):
def __init__(self, input_size: int, hidden_size: int, plasticity_rule: Callable, tau_trace: float,
dynamics: NeuronModel) -> None:
super().__init__()
self.input_size = input_size
self.hidden_size = hidden_size
self.plasticity_rule = plasticity_rule
self.dynamics = dynamics
self.decay_trace = math.exp(-1.0 / tau_trace)
self.W = torch.nn.Parameter(torch.Tensor(hidden_size + hidden_size, input_size))
self.reset_parameters()
def forward(self, x: torch.Tensor, mem: Optional[torch.Tensor] = None, states: Optional[Tuple[
List[torch.Tensor], List[torch.Tensor], torch.Tensor, torch.Tensor]] = None) -> Tuple[
torch.Tensor, torch.Tensor, torch.Tensor, List[torch.Tensor]]:
batch_size, sequence_length, _ = x.size()
if states is None:
key_states = self.dynamics.initial_states(batch_size, self.hidden_size, x.dtype, x.device)
val_states = self.dynamics.initial_states(batch_size, self.hidden_size, x.dtype, x.device)
key_trace = torch.zeros(batch_size, self.hidden_size, dtype=x.dtype, device=x.device)
val_trace = torch.zeros(batch_size, self.hidden_size, dtype=x.dtype, device=x.device)
else:
key_states, val_states, key_trace, val_trace = states
if mem is None:
mem = torch.zeros(batch_size, self.hidden_size, self.hidden_size, dtype=x.dtype, device=x.device)
i = torch.nn.functional.linear(x, self.W)
ik, iv = i.chunk(2, dim=2)
key_output_sequence = []
val_output_sequence = []
for t in range(sequence_length):
# Key-layer
key, key_states = self.dynamics(ik.select(1, t), key_states)
# Current from key-layer to value-layer ('bij,bj->bi', mem, key)
ikv_t = 0.2 * (key.unsqueeze(1) * mem).sum(2)
# Value-layer
val, val_states = self.dynamics(iv.select(1, t) + ikv_t, val_states)
# Update traces
key_trace = exp_convolve(key, key_trace, self.decay_trace)
val_trace = exp_convolve(val, val_trace, self.decay_trace)
# Update memory
delta_mem = self.plasticity_rule(key_trace, val_trace, mem)
mem = mem + delta_mem
key_output_sequence.append(key)
val_output_sequence.append(val)
states = [key_states, val_states, key_trace, val_trace]
return mem, torch.stack(key_output_sequence, dim=1), torch.stack(val_output_sequence, dim=1), states
def reset_parameters(self) -> None:
torch.nn.init.xavier_uniform_(self.W, gain=math.sqrt(2))
class WritingLayerReLU(torch.nn.Module):
def __init__(self, input_size: int, hidden_size: int, plasticity_rule: Callable) -> None:
super().__init__()
self.input_size = input_size
self.hidden_size = hidden_size
self.plasticity_rule = plasticity_rule
self.W = torch.nn.Parameter(torch.Tensor(hidden_size + hidden_size, input_size))
self.reset_parameters()
def forward(self, x: torch.Tensor, states: Optional[torch.Tensor] = None) -> Tuple[
torch.Tensor, torch.Tensor, torch.Tensor]:
batch_size, sequence_length, _ = x.size()
if states is None:
mem = torch.zeros(batch_size, self.hidden_size, self.hidden_size, dtype=x.dtype, device=x.device)
else:
mem = states
i = torch.nn.functional.linear(x, self.W)
ik, iv = i.chunk(2, dim=2)
key_output_sequence = []
val_output_sequence = []
for t in range(sequence_length):
# Key-layer
key = torch.nn.functional.relu(ik.select(1, t))
# Value-layer
val = torch.nn.functional.relu(iv.select(1, t))
# Update memory
delta_mem = self.plasticity_rule(key, val, mem)
mem = mem + delta_mem
key_output_sequence.append(key)
val_output_sequence.append(val)
return mem, torch.stack(key_output_sequence, dim=1), torch.stack(val_output_sequence, dim=1)
def reset_parameters(self) -> None:
torch.nn.init.xavier_uniform_(self.W, gain=math.sqrt(2))