| 10 | |
| 11 | |
| 12 | class MyTest_Graves(unittest.TestCase): |
| 13 | def test_1(self): |
| 14 | batch_size = 10 |
| 15 | seq_len = 20 |
| 16 | seq_len_fs = 15 |
| 17 | num_chars = 7 |
| 18 | dim_x = 5 |
| 19 | dim_c = 6 |
| 20 | dim_p = 9 |
| 21 | dim_z = 7 |
| 22 | dim_latent = 10 |
| 23 | dim_f = 13 |
| 24 | valid_lens_fs = [random.randint(1, seq_len_fs) for _ in range(batch_size)] |
| 25 | |
| 26 | param_graves = ParamGraves(dim_x=dim_x, dim_c=dim_c, dim_p=dim_p, dim_z=dim_z) |
| 27 | param_vrnn = ParamVRNN(param_graves=param_graves, dim_latent=dim_latent, dim_f=dim_f) |
| 28 | |
| 29 | net = NetworkVRNN(param_vrnn) |
| 30 | |
| 31 | xs = torch.randn(seq_len, batch_size, dim_x) |
| 32 | cs = torch.randn(num_chars, batch_size, dim_c) |
| 33 | fs = torch.randn(seq_len_fs, batch_size, dim_f) |
| 34 | init_h = net.get_zero_hidden_states(batch_size=batch_size, device=xs.device) |
| 35 | |
| 36 | # forward |
| 37 | out_dict = net(xs=xs, cs=cs, fs=fs, init_hidden_states=init_h, valid_lens_fs=valid_lens_fs) |
| 38 | |
| 39 | # check ps |
| 40 | ps = out_dict['ps'] |
| 41 | assert len(ps.shape) == 3 |
| 42 | assert ps.size(0) == seq_len |
| 43 | assert ps.size(1) == batch_size |
| 44 | assert ps.size(2) == dim_p |
| 45 | |
| 46 | # check attn_cs |
| 47 | attn_cs = out_dict['attn_cs'] |
| 48 | assert len(attn_cs.shape) == 3 |
| 49 | assert attn_cs.size(0) == seq_len |
| 50 | assert attn_cs.size(1) == batch_size |
| 51 | assert attn_cs.size(2) == dim_c |
| 52 | |
| 53 | # check attn_weights |
| 54 | attn_weights = out_dict['attn_weights'] |
| 55 | assert len(attn_weights.shape) == 3 |
| 56 | assert attn_weights.size(0) == seq_len |
| 57 | assert attn_weights.size(1) == batch_size |
| 58 | assert attn_weights.size(2) == num_chars + 1 |
| 59 | |
| 60 | # check attn_hs |
| 61 | attn_hs = out_dict['attn_hs'] |
| 62 | assert len(attn_hs.shape) == 3 |
| 63 | assert attn_hs.size(0) == seq_len |
| 64 | assert attn_hs.size(1) == batch_size |
| 65 | assert attn_hs.size(2) == net.dim_attn_rnn_h |
| 66 | |
| 67 | # check decode_hs |
| 68 | decode_hs = out_dict['decode_hs'] |
| 69 | assert len(decode_hs.shape) == 3 |
nothing calls this directly
no outgoing calls
no test coverage detected