(self, xWxr_t, xWxz_t, xWxh_t, is_start, h_t1, h0)
| 98 | return h |
| 99 | |
| 100 | def recurrence(self, xWxr_t, xWxz_t, xWxh_t, is_start, h_t1, h0): |
| 101 | h_t = T.switch( |
| 102 | T.eq(is_start, 1), |
| 103 | self.get_ht(xWxr_t, xWxz_t, xWxh_t, h0), |
| 104 | self.get_ht(xWxr_t, xWxz_t, xWxh_t, h_t1) |
| 105 | ) |
| 106 | return h_t |
| 107 | |
| 108 | def output(self, Xflat, startPoints): |
| 109 | # Xflat should be (NT, D) |