-
Notifications
You must be signed in to change notification settings - Fork 1
Expand file tree
/
Copy pathevaluate.py
More file actions
57 lines (43 loc) · 2.05 KB
/
Copy pathevaluate.py
File metadata and controls
57 lines (43 loc) · 2.05 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
import torch
from utils import *
from load import loadTestData
from constants import BOS_IDX, DEC_MAX_LENGTH
#import visualize_online
USE_CUDA = torch.cuda.is_available()
device = torch.device("cuda" if USE_CUDA else "cpu")
def evaluateTestFile(decoder, spine, mask_generator, voc_dec, test_data):
trg_embs, ctx_embs, trg_words, def_sents, ctx_sents = test_data
out_file = open('outfile.txt', 'w')
bs = 1024
max_sum = 0
for i in range(0, len(trg_embs), bs):
trg_emb = torch.FloatTensor(trg_embs[i: i+bs]).to(device)
ctx_emb = torch.FloatTensor(ctx_embs[i: i+bs]).to(device)
sp_z, sp_w, loss_terms = spine(trg_emb)
aligned_ctx, sense_vec, attn, _ = mask_generator(sp_z, sp_w, ctx_emb)
max_sum += torch.sum(torch.max(attn, dim=1)[0])
decoder_input = torch.LongTensor([[BOS_IDX] * len(trg_emb)]).to(device)
decoder_hidden = trg_emb.unsqueeze(0)
decoder_hidden2 = aligned_ctx.unsqueeze(0)
for j in range(DEC_MAX_LENGTH):
decoder_output, decoder_hidden, decoder_hidden2 = decoder(
decoder_input, decoder_hidden, decoder_hidden2, sense_vec.unsqueeze(0)
)
# decoder_output: (bs, n_voc), decoder_input: (1, bs)
out = torch.argmax(decoder_output, dim=1) # (bs,)
preds = torch.cat((preds, out.unsqueeze(1)), dim=1) if j else out.unsqueeze(1)
decoder_input = out.unsqueeze(0)
# preds: (bs, max_len)
for j in range(preds.shape[0]):
w = trg_words[i+j]
ctx = ctx_sents[i+j]
truth = def_sents[i+j]
pre = id2sent(preds[j], voc_dec)
out_file.write('{} ; {} ; {} ; {}\n'.format(w, ctx, truth, pre))
print('avg highest attn value: {:.3f}'.format(max_sum.item()/len(trg_embs)))
out_file.close()
def runTest(args):
voc_dec, test_data = loadTestData(args)
torch.set_grad_enabled(False)
decoder, spine, mask_generator = load_model(args)
evaluateTestFile(decoder, spine, mask_generator, voc_dec, test_data)