| 12 | return np.random.rand(*args) * (b - a) + a |
| 13 | |
| 14 | class LstmParam: |
| 15 | def __init__(self, mem_cell_ct, x_dim): |
| 16 | self.mem_cell_ct = mem_cell_ct |
| 17 | self.x_dim = x_dim |
| 18 | concat_len = x_dim + mem_cell_ct |
| 19 | # weight matrices |
| 20 | self.wg = rand_arr(-0.1, 0.1, mem_cell_ct, concat_len) |
| 21 | self.wi = rand_arr(-0.1, 0.1, mem_cell_ct, concat_len) |
| 22 | self.wf = rand_arr(-0.1, 0.1, mem_cell_ct, concat_len) |
| 23 | self.wo = rand_arr(-0.1, 0.1, mem_cell_ct, concat_len) |
| 24 | # bias terms |
| 25 | self.bg = rand_arr(-0.1, 0.1, mem_cell_ct) |
| 26 | self.bi = rand_arr(-0.1, 0.1, mem_cell_ct) |
| 27 | self.bf = rand_arr(-0.1, 0.1, mem_cell_ct) |
| 28 | self.bo = rand_arr(-0.1, 0.1, mem_cell_ct) |
| 29 | # diffs (derivative of loss function w.r.t. all parameters) |
| 30 | self.wg_diff = np.zeros((mem_cell_ct, concat_len)) |
| 31 | self.wi_diff = np.zeros((mem_cell_ct, concat_len)) |
| 32 | self.wf_diff = np.zeros((mem_cell_ct, concat_len)) |
| 33 | self.wo_diff = np.zeros((mem_cell_ct, concat_len)) |
| 34 | self.bg_diff = np.zeros(mem_cell_ct) |
| 35 | self.bi_diff = np.zeros(mem_cell_ct) |
| 36 | self.bf_diff = np.zeros(mem_cell_ct) |
| 37 | self.bo_diff = np.zeros(mem_cell_ct) |
| 38 | |
| 39 | def apply_diff(self, lr = 1): |
| 40 | self.wg -= lr * self.wg_diff |
| 41 | self.wi -= lr * self.wi_diff |
| 42 | self.wf -= lr * self.wf_diff |
| 43 | self.wo -= lr * self.wo_diff |
| 44 | self.bg -= lr * self.bg_diff |
| 45 | self.bi -= lr * self.bi_diff |
| 46 | self.bf -= lr * self.bf_diff |
| 47 | self.bo -= lr * self.bo_diff |
| 48 | # reset diffs to zero |
| 49 | self.wg_diff = np.zeros_like(self.wg) |
| 50 | self.wi_diff = np.zeros_like(self.wi) |
| 51 | self.wf_diff = np.zeros_like(self.wf) |
| 52 | self.wo_diff = np.zeros_like(self.wo) |
| 53 | self.bg_diff = np.zeros_like(self.bg) |
| 54 | self.bi_diff = np.zeros_like(self.bi) |
| 55 | self.bf_diff = np.zeros_like(self.bf) |
| 56 | self.bo_diff = np.zeros_like(self.bo) |
| 57 | |
| 58 | class LstmState: |
| 59 | def __init__(self, mem_cell_ct, x_dim): |