MCPcopy Create free account
hub / github.com/Hunyuan-PromptEnhancer/PromptEnhancer / predict

Method predict

inference/app.py:49–111  ·  view source on GitHub ↗
(
        self,
        prompt_cot,
        sys_prompt="请根据用户的输入,生成思考过程的思维链并改写提示词:",
        temperature=0.0,
        top_p=1.0,
        max_new_tokens=2048,
        device="cuda",
    )

Source from the content-addressed store, hash-verified

47
48 @torch.inference_mode()
49 def predict(
50 self,
51 prompt_cot,
52 sys_prompt="请根据用户的输入,生成思考过程的思维链并改写提示词:",
53 temperature=0.0,
54 top_p=1.0,
55 max_new_tokens=2048,
56 device="cuda",
57 ):
58 org_prompt_cot = prompt_cot
59 try:
60 user_prompt_format = sys_prompt + "\n" + org_prompt_cot
61 messages = [
62 {
63 "role": "user",
64 "content": [
65 {"type": "text", "text": user_prompt_format},
66 ],
67 }
68 ]
69
70 text = self.processor.apply_chat_template(
71 messages, tokenize=False, add_generation_prompt=True
72 )
73 image_inputs, video_inputs = process_vision_info(messages)
74 inputs = self.processor(
75 text=[text],
76 images=image_inputs,
77 videos=video_inputs,
78 padding=True,
79 return_tensors="pt",
80 )
81 inputs = inputs.to(device)
82
83 # 注意:原始代码固定 do_sample=False,top_k=5, top_p=0.9,这里保持一致
84 generated_ids = self.model.generate(
85 **inputs,
86 max_new_tokens=2048, # 与原始代码保持一致(未使用 max_new_tokens 参数)
87 temperature=float(temperature),
88 do_sample=False,
89 top_k=5,
90 top_p=0.9
91 )
92 generated_ids_trimmed = [
93 out_ids[len(in_ids):]
94 for in_ids, out_ids in zip(inputs.input_ids, generated_ids)
95 ]
96 output_text = self.processor.batch_decode(
97 generated_ids_trimmed,
98 skip_special_tokens=True,
99 clean_up_tokenization_spaces=False,
100 )
101 output_res = output_text[0]
102 assert output_res.count("think>") == 2
103 prompt_cot = output_res.split("think>")[-1]
104 if prompt_cot.startswith("\n"):
105 prompt_cot = prompt_cot[1:]
106 prompt_cot = replace_single_quotes(prompt_cot)

Callers 2

run_singleFunction · 0.45
run_batchFunction · 0.45

Calls 2

process_vision_infoFunction · 0.70
replace_single_quotesFunction · 0.70

Tested by

no test coverage detected