File size: 4,192 Bytes
4ca4e4c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
import torch
from data_utils import get_eval_dataset
import numpy as np


def get_score(args, tokenizer, x, pred, mask, i):
    x = x[i]
    pred = pred[i]
    mask = mask[i]

    # str_acc = int(torch.equal(x * mask, pred * mask)) # I don't think I know what equal does
    str_acc = int(sum(torch.eq(x * mask, pred * mask)) == sum(mask))
    
    if sum(mask) == 0:
        char_acc = 1
    else:
        # char_acc = sum([m * int(c1 == c2) for (c1, c2, m) in zip(x, pred, mask)]) / sum(mask)
        char_acc = sum(mask * torch.eq(x, pred)) / sum(mask)

    return str_acc, char_acc


def evaluation(args, model, tokenizer):
    if args.eval_jump_type == 'exponential':
        lengths = np.unique(np.logspace(np.log10(args.min_eval_length), np.log10(args.max_eval_length), num=args.eval_exp_num_jumps, dtype=int))
    if args.eval_jump_type == 'linear':
        lengths = np.arange(args.min_eval_length, args.max_eval_length, args.eval_linear_jump_size)
    # lengths = np.arange(args.min_eval_length, args.max_eval_length)
    
    str_acc_mean_list = []
    str_acc_std_list = []
    char_accuracy_list = []
    if args.print:
        print("\n")
    
    for length in lengths:
        str_acc_batch = np.zeros(args.eval_num_batches)
        char_acc_mean = 0

        for j in range(args.eval_num_batches):
            long_dataset = get_eval_dataset(args, tokenizer, length, length)     
            batch = next(iter(long_dataset))
            
            x = batch['input_ids'].to('cuda')
            y = batch['output_ids'].to('cuda')
            mask = batch['mask'].to('cuda')

            attention_mask = torch.ones((x.shape[1], x.shape[1]))
            attention_mask = (torch.triu(attention_mask, diagonal=0) - torch.triu(attention_mask, diagonal=args.window)).T.to('cuda')
            
            with torch.no_grad():
                # prediction
                if args.model=="lstm":
                    assert False # Not tested
                    state = model.init_hidden(args.eval_batch_size, 'cuda')
                    logits, state = model(x, state)
                # elif args.model=="mamba":
                #     logits = model(x)[0]
                else:
                    logits = model(x, attention_mask=attention_mask, return_dict=True)['logits']
                
                # greedy decoding
                pred = torch.argmax(logits, dim=-1)

                # evaluation
                for i in range(len(x)):
                    str_acc, char_acc = get_score(args, tokenizer, y, pred, mask, i) 

                    str_acc_batch[j] += str_acc
                    char_acc_mean    += char_acc

        if args.print:
            print("v"*100)
            print("COMPLETE EXAMPLE: ", tokenizer.to_string(batch['input'][0], pytorch=False))
            # print("EXAMPLE:", batch['input'][0])
            print("-"*100)
            print("INPUT EXAMPLE: ", tokenizer.to_string(batch['input_ids'][0][batch['mask'][0]==1]))
            print("OUTPUT EXAMPLE:", tokenizer.to_string(batch['output_ids'][0][batch['mask'][0]==1]))
            # print("TOKENIZED:", batch['input_ids'][0][batch['mask'][0]==1])
            print("-"*100)
            print("PREDICTION:    ", tokenizer.to_string(pred[0][batch['mask'][0]==1]))
            # print("PREDICTION:", pred[-1][batch['mask'][0]==1])
            print("^"*100)
        

        str_acc_batch = str_acc_batch/args.eval_batch_size
        # str_acc_batch = str_acc_batch/len(x)
        mean_str_acc = float(np.mean(str_acc_batch))
        std_str_acc = float(np.std(str_acc_batch))

        str_acc_mean_list.append(mean_str_acc)
        str_acc_std_list.append(std_str_acc)

        mean_char_acc = char_acc_mean/(args.eval_batch_size*args.eval_num_batches)
        # mean_char_acc = char_acc_mean/(len(x)*args.eval_num_batches)
        char_accuracy_list.append(mean_char_acc.item())
        
        # print(f"{args.eval_task}; len {length}: {mean_str_acc} +- {std_str_acc}; char: {mean_char_acc}")
        print(f"{args.eval_task};len {length};char: {mean_char_acc}")
    if args.print:
        print("\n")        
    
    return str_acc_mean_list, str_acc_std_list, char_accuracy_list