MCPcopy Create free account
hub / github.com/OrangeInSouth/DeePEn / MainModelThread

Class MainModelThread

src/main_model_thread.py:11–102  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

9
10
11class MainModelThread(threading.Thread):
12 def __init__(self, main_model, main_model_tokenizer, assist_model_tokenizer, information_dict,
13 learning_rate, learning_epochs_nums, result_save_dir,
14 ensemble_model_output_ids_queue,
15 assist_model_score_queue_list, main_model_probability_transfer_matrix_list,
16 assist_model_probability_transfer_matrix_list, device, device_compute, early_stop_string_list=None):
17 self.model = main_model
18 self.tokenizer = main_model_tokenizer
19 self.assist_model_tokenizer = assist_model_tokenizer
20 self.information_dict = information_dict
21 self.model_streamer = TextStreamer(self.tokenizer)
22 self.learning_rate = learning_rate
23 self.learning_epochs_nums = learning_epochs_nums
24 self.result_save_dir = result_save_dir
25 self.ensemble_model_output_ids_queue = ensemble_model_output_ids_queue
26 self.assist_model_score_queue_list = assist_model_score_queue_list
27 self.main_model_probability_transfer_matrix_list = main_model_probability_transfer_matrix_list
28 self.assist_model_probability_transfer_matrix_list = assist_model_probability_transfer_matrix_list
29 self.device = device
30 self.device_compute = device_compute
31 self.early_stop_string_list = early_stop_string_list
32
33 super().__init__()
34
35 def run(self) -> None:
36 main_model_logits_processor_list = LogitsProcessorList()
37
38 processor_factory = ModelProcessorFactory()
39
40 # 传递其他参数
41 additional_kwargs = {
42 "learning_rate": self.learning_rate,
43 "ensemble_weight": self.information_dict['ensemble_weight'],
44 "learning_epochs_nums": self.learning_epochs_nums,
45 "ensemble_model_output_ids_queue": self.ensemble_model_output_ids_queue,
46 "assist_model_score_queue_list": self.assist_model_score_queue_list,
47
48 "main_model_probability_transfer_matrix_list": self.main_model_probability_transfer_matrix_list,
49 "assist_model_probability_transfer_matrix_list": self.assist_model_probability_transfer_matrix_list,
50 "result_save_dir": self.result_save_dir,
51 "main_model_tokenizer": self.tokenizer,
52 "assist_model_tokenizer": self.assist_model_tokenizer,
53 "device": self.device,
54 "device_compute": self.device_compute,
55 "early_stop_string_list": self.early_stop_string_list,
56 }
57 # 创建对象
58 logits_processor_mode = self.information_dict['logits_processor_mode']
59 logits_processor_instance = processor_factory.create_processor(logits_processor_mode,
60 **additional_kwargs)
61 main_model_logits_processor_list.append(logits_processor_instance)
62 # main_model_logits_processor_list.append()
63
64 main_model_input = self.information_dict['main_model_input']
65 max_new_tokens = self.information_dict['max_new_tokens']
66 main_model_input_ids = self.tokenizer(main_model_input, return_tensors="pt",
67 add_special_tokens=False).input_ids.to(self.device)
68 generation_kwargs = {

Callers 2

mainFunction · 0.90
mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected