File size: 3,370 Bytes
31dc8dc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from accelerate.logging import get_logger
logger = get_logger(__name__, log_level="INFO")


import torch
class UniversalPrompting():
    def __init__(self, text_tokenizer,
                 max_prompt_len=8000, max_gen_length=377, ignore_id=-100, system_prompt=None):
        """
        :param text_tokenizer: original text tokenizer
        """
        self.text_tokenizer = text_tokenizer
        self.max_gen_length = max_gen_length
        self.max_prompt_len = max_prompt_len

    def lm_prompt(self, text_ids_pairs):
        prompts_list, responses_list = text_ids_pairs
        pad_id = self.text_tokenizer.pad_token_id

        if responses_list.shape[1] < self.max_gen_length:
            max_seq_len = prompts_list.shape[1] + responses_list.shape[1]
        else:
            max_seq_len = prompts_list.shape[1] + self.max_gen_length

        sequence_ids = []
        attention_masks = []
        label_ids = []

        for prompt_ids, resp_ids in zip(prompts_list, responses_list):
            prompt_ids = prompt_ids.tolist()
            resp_ids   = resp_ids.tolist()

            temp_ids = prompt_ids + resp_ids
            temp_masks = [1] * len(temp_ids)
            temp_labels = temp_ids.copy()

            if len(temp_ids) < max_seq_len:
                pad_len = max_seq_len - len(temp_ids)
                temp_ids.extend([pad_id] * pad_len)
                temp_labels.extend([pad_id] * pad_len)
                temp_masks.extend([0] * pad_len)
            else:
                temp_ids = temp_ids[:max_seq_len]
                temp_labels = temp_labels[:max_seq_len]
                temp_masks = temp_masks[:max_seq_len]

            sequence_ids.append(torch.tensor(temp_ids).unsqueeze(0))
            attention_masks.append(torch.tensor(temp_masks).unsqueeze(0))
            label_ids.append(torch.tensor(temp_labels).unsqueeze(0))

        input_ids = torch.cat(sequence_ids, dim=0)
        attention_masks = torch.cat(attention_masks, dim=0)
        label_ids = torch.cat(label_ids, dim=0)
        

        return input_ids, label_ids, prompts_list.shape[1]

        
    

    def mask_prompt(self):
        pass

    def __call__(self, input):
        prompts, responses = input

        enc = self.text_tokenizer(
            prompts,
            padding=False,
            truncation=False,
            return_length=True
        )
        lengths = enc["length"]
        keep_indices = [i for i, L in enumerate(lengths) if L <= self.max_prompt_len]
        drop_num = len(prompts) - len(keep_indices)
        
        prompts = [self.text_tokenizer.apply_chat_template(
            [{"role": "user", "content": prompts[i]}],
            tokenize=False,
            add_generation_prompt=True
        ) for i in keep_indices]
        responses = [responses[i] for i in keep_indices]

        prompt_ids = self.text_tokenizer(
            prompts,
            padding=True,
            return_tensors="pt",
            padding_side = "left"
        )['input_ids']
        response_ids = self.text_tokenizer(
            responses,
            padding=True,
            return_tensors="pt",
            padding_side = "right"
        )['input_ids']

        input_ids_lm, labels_lm, start_pos = self.lm_prompt((prompt_ids, response_ids))
        return input_ids_lm, labels_lm, start_pos, drop_num


if __name__ == '__main__':
    pass