Views
No views yet
1from typing import List, TypedDict
2from dataclasses import dataclass
3from itertools import chain
4
5from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
6import torch
7
8
9@dataclass
10class H2PersonaChatHyperparametersV1:
11 """
12 chat_history_pair_length: int - количество пар диалога с конца
13 """
14
15 model_name: str = "facebook/bart-base"
16 chat_history_pair_length: int = 7
17
18 persona_max_length: int = 14
19 chat_max_length: int = 25
20
21 debug_status: int = 0
22
23
24class PersonaChatDatasetSampleV1(TypedDict):
25 """
26 persona: List[str] - набор предложений фактов персоны
27 history: List[str] - набор предложений истории переписки
28 """
29
30 persona: List[str]
31 history: List[str]
32 sample_id: str
33
34
35class H2Seq2SeqInferenceSampleDictV1(TypedDict):
36 input_ids: List[int]
37 attention_mask: List[int]
38
39
40class H2Seq2SeqInferenceSampleDictV2(TypedDict):
41 input_ids: torch.Tensor
42 attention_mask: torch.Tensor
43
44
45def flat_list(list_of_lists: List[List]) -> List:
46 return list(chain.from_iterable(list_of_lists))
47
48
49class H2Seq2SeqInferencePersonaSampleV1:
50 def __init__(
51 self,
52 dataset_sample: PersonaChatDatasetSampleV1,
53 tokenizer: AutoTokenizer,
54 hyperparameters: H2PersonaChatHyperparametersV1,
55 ) -> None:
56 self.dataset_sample = dataset_sample
57 self.tokenizer = tokenizer
58 self.hyperparameters = hyperparameters
59
60 def add_spaces_after(
61 self,
62 items: List[str],
63 ) -> List[str]:
64 items = [item + " " for item in items]
65 return items
66
67 @property
68 def bos_token_id(self):
69 if "t5" in self.hyperparameters.model_name:
70 return []
71
72 if self.tokenizer.bos_token_id is None:
73 return []
74
75 return [self.tokenizer.bos_token_id]
76
77 @property
78 def eos_token_id(self):
79 if self.tokenizer.eos_token_id is None:
80 return []
81
82 return [self.tokenizer.eos_token_id]
83
84 def add_sep_beetween(self, items: List[str], sep=" EOS ") -> List[str]:
85 for i in range(1, len(items)):
86 items[i] = sep + items[i]
87
88 return items
89
90 def add_spaces_between(self, items: List[str]) -> List[str]:
91 items = self.add_spaces_after(items)
92 items[-1] = items[-1].strip()
93 return items
94
95 def get_sample(self) -> H2Seq2SeqInferenceSampleDictV1:
96
97 dialog_history = self.dataset_sample["history"]
98 dialog_history = dialog_history[-self.hyperparameters.chat_history_pair_length * 2 - 1 :]
99 dialog_history = self.add_sep_beetween(dialog_history)
100
101 persona = self.dataset_sample["persona"]
102 persona = self.add_sep_beetween(
103 persona,
104 sep=" ",
105 )
106
107 KNOWLEDGE_IDS = self.tokenizer.encode(
108 " [KNOWLEDGE] ",
109 add_special_tokens=False,
110 )
111 CONTEXT_IDS = self.tokenizer.encode(
112 " [CONTEXT]",
113 add_special_tokens=False,
114 )
115
116 encoded_history = self.tokenizer.batch_encode_plus(
117 dialog_history,
118 add_special_tokens=False,
119 truncation=True,
120 max_length=self.hyperparameters.chat_max_length,
121 )
122 encoded_history = flat_list(encoded_history["input_ids"])
123
124 encoded_persona = self.tokenizer.batch_encode_plus(
125 persona,
126 add_special_tokens=False,
127 truncation=True,
128 max_length=self.hyperparameters.persona_max_length,
129 )
130
131 encoded_persona = flat_list(encoded_persona["input_ids"])
132
133 input_ids = [
134 *self.bos_token_id,
135 *CONTEXT_IDS,
136 *encoded_history,
137 *KNOWLEDGE_IDS,
138 *encoded_persona,
139 *self.eos_token_id,
140 ]
141
142 attention_mask = [1] * len(input_ids)
143
144 return H2Seq2SeqInferenceSampleDictV1(
145 input_ids=input_ids,
146 attention_mask=attention_mask,
147 )
148
149
150class DialogBotV1:
151 def __init__(
152 self,
153 model: AutoModelForSeq2SeqLM,
154 tokenizer: AutoTokenizer,
155 hyperparameters: H2PersonaChatHyperparametersV1,
156 history: List[str] = None,
157 persona: List[str] = None,
158 device: str = "cuda",
159 shuffle_persona: bool = True,
160 ):
161 self.model = model
162
163 self.tokenizer = tokenizer
164 self.hyperparameters = hyperparameters
165 self.device = device
166 self.shuffle_persona = shuffle_persona
167
168 self.debug_status = hyperparameters.debug_status
169
170 if history is None:
171 self.history = []
172 self.history = history
173
174 if persona is None:
175 self.persona = []
176 self.persona = persona
177
178 def _get_sample(
179 self,
180 persona: List[str],
181 history: List[str],
182 ) -> H2Seq2SeqInferenceSampleDictV1:
183 dataset_sample = PersonaChatDatasetSampleV1(
184 persona=persona,
185 history=history,
186 )
187
188 sample = H2Seq2SeqInferencePersonaSampleV1(
189 tokenizer=self.tokenizer,
190 hyperparameters=self.hyperparameters,
191 dataset_sample=dataset_sample,
192 )
193 sample = sample.get_sample()
194 print(self.tokenizer.decode(sample['input_ids']))
195
196 for key in sample.keys():
197 sample[key] = torch.tensor(sample[key]).unsqueeze(0).to(self.device)
198
199 return sample
200
201 def next_response(
202 self,
203 **generation_params,
204 ) -> str:
205 """
206 делает предсказание на основе текущей истории
207 и персоны
208 """
209
210 sample = self._get_sample(
211 persona=self.persona,
212 history=self.history,
213 )
214 answer = self.generate_response(
215 sample,
216 **generation_params,
217 )
218 answer = self.tokenizer.batch_decode(
219 answer,
220 skip_special_tokens=True,
221 )
222 self.history.append(answer[0])
223 return answer[0]
224
225 def generate_response(
226 self,
227 sample: H2Seq2SeqInferenceSampleDictV1,
228 **generation_params,
229 ):
230 """
231 generation_params - https://huggingface.co/docs/transformers/v4.24.0/en/main_classes/text_generation
232 """
233 with torch.no_grad():
234 return self.model.generate(
235 **sample,
236 **generation_params,
237 )
238
239
240# facebook/mbart-large-50
241PRETRAINED_MODEL_NAME_OR_PATH = "DeepPavlov/mbart-large-50-ru-persona-chat"
242
243PAIR_DIALOG_HISTORY_LENGTH = 2
244
245# CHAT_MAX_LENGTH for single sentence
246CHAT_MAX_LENGTH = 25
247# PERSONA_MAX_LENGTH for single sentence
248PERSONA_MAX_LENGTH = 19
249
250device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
251model = AutoModelForSeq2SeqLM.from_pretrained(PRETRAINED_MODEL_NAME_OR_PATH)
252model.to(device)
253model.eval()
254
255tokenizer = AutoTokenizer.from_pretrained(PRETRAINED_MODEL_NAME_OR_PATH)
256
257if torch.cuda.is_available():
258 model.half()
259
260hyperparameters = H2PersonaChatHyperparametersV1(
261 chat_history_pair_length=PAIR_DIALOG_HISTORY_LENGTH,
262 persona_max_length=PERSONA_MAX_LENGTH,
263 chat_max_length=CHAT_MAX_LENGTH,
264 model_name=PRETRAINED_MODEL_NAME_OR_PATH,
265)
266
267
268persona = [
269 "Я люблю играть с милыми песиками",
270 "Я ненавижу лук и броколли"
271]
272
273history = [
274 "Привет. Ты любишь лук?"
275]
276
277persona_bot = DialogBotV1(
278 model=model,
279 tokenizer=tokenizer,
280 hyperparameters=hyperparameters,
281 history=history,
282 persona=persona,
283 device=device,
284 )
285
286GENERATION_PARAMS = {
287 "max_new_tokens": 60,
288 "penalty_alpha": 0.15,
289 "top_k": 10
290}
291response = persona_bot.next_response(
292 **GENERATION_PARAMS,
293)
294
295print(response)
2961def ru_persona_chat_dataset_tranformer_v1(
2 initial_dataset_path: str,
3 output_folder: str,
4) -> None:
5 """
6 example
7 ru_persona_chat_dataset_tranformer_v1(
8 initial_dataset_path="./datasets/ru_persona_chat/dialogues.tsv",
9 output_folder="./datasets/ru_persona_chat",
10 )
11 """
12 assert initial_dataset_path is not None, "initial_dataset_path is None"
13 assert output_folder is not None, "output_folder is None"
14
15 dataset = pd.read_csv(initial_dataset_path, sep="\t")
16 split_ratio = int(len(dataset) * 0.95)
17 train_dataset = dataset[:split_ratio]
18 valid_dataset = dataset[split_ratio:]
19
20 print(f"Dataset lengths: train {len(train_dataset)}, valid {len(valid_dataset)}")
21 # save csv files
22 train_dataset.to_csv(output_folder + "/train.csv", index=False)
23 valid_dataset.to_csv(output_folder + "/valid.csv", index=False)
24 print("Datasets saved.")