(self, xWxr_t, xWxz_t, xWxh_t, is_start, h_t1, h0)
| 52 | return h |
| 53 | |
| 54 | def recurrence(self, xWxr_t, xWxz_t, xWxh_t, is_start, h_t1, h0): |
| 55 | h_t = T.switch( |
| 56 | T.eq(is_start, 1), |
| 57 | self.get_ht(xWxr_t, xWxz_t, xWxh_t, h0), |
| 58 | self.get_ht(xWxr_t, xWxz_t, xWxh_t, h_t1) |
| 59 | ) |
| 60 | return h_t |
| 61 | |
| 62 | def output(self, Xflat, startPoints): |
| 63 | # Xflat should be (NT, D) |