(self, first=1, nchars=0, stop=-1)
| 70 | |
| 71 | |
| 72 | def sample(self, first=1, nchars=0, stop=-1): |
| 73 | res = [first] |
| 74 | dy.renew_cg() |
| 75 | state = self.builder.initial_state() |
| 76 | |
| 77 | R = dy.parameter(self.R) |
| 78 | bias = dy.parameter(self.bias) |
| 79 | cw = first |
| 80 | while True: |
| 81 | x_t = dy.lookup(self.lookup, cw) |
| 82 | state = state.add_input(x_t) |
| 83 | y_t = state.output() |
| 84 | r_t = bias + (R * y_t) |
| 85 | ydist = dy.softmax(r_t) |
| 86 | dist = ydist.vec_value() |
| 87 | rnd = random.random() |
| 88 | for i,p in enumerate(dist): |
| 89 | rnd -= p |
| 90 | if rnd <= 0: break |
| 91 | res.append(i) |
| 92 | cw = i |
| 93 | if cw == stop: break |
| 94 | if nchars and len(res) > nchars: break |
| 95 | return res |
| 96 | |
| 97 | if __name__ == '__main__': |
| 98 | parser = argparse.ArgumentParser() |
no test coverage detected