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

Method run

src/main_model_thread.py:35–102  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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 = {
69 "input_ids": main_model_input_ids,
70 "max_new_tokens": max_new_tokens,
71 "do_sample": False,
72 "num_beams": 1,
73 "eos_token_id": self.tokenizer.eos_token_id,
74 "bos_token_id": self.tokenizer.bos_token_id,
75 # "pad_token_id": self.tokenizer.pad_token_id
76 }
77
78 # generate_ids = self.model.generate(**generation_kwargs,pad_token_id=self.tokenizer.eos_token_id,
79 # logits_processor=main_model_logits_processor_list,
80 # streamer=self.model_streamer)
81 generate_ids = self.model.generate(**generation_kwargs, pad_token_id=self.tokenizer.eos_token_id,
82 logits_processor=main_model_logits_processor_list)
83
84 text = self.tokenizer.decode(generate_ids[0])
85 # print(text)
86 result_process_parameter = self.information_dict['result_process_parameter']
87 split_key_before_list = result_process_parameter["split_key_before"]
88 split_key_behind_list = result_process_parameter["split_key_behind"]
89
90 model_answer, prediction = answer_extract(text, self.information_dict['demon_count'], split_key_before_list,
91 split_key_behind_list)
92 print(self.information_dict['question'])

Callers

nothing calls this directly

Calls 3

create_processorMethod · 0.95
answer_extractFunction · 0.90

Tested by

no test coverage detected