Skip to content

Commit f212845

Browse files
authored
Migrate 4 operators written by @yunwenkai2003 from DataFlow421 (OpenDCAI#202)
1 parent 133499d commit f212845

4 files changed

Lines changed: 463 additions & 0 deletions

File tree

Lines changed: 78 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1 +1,79 @@
11
from dataflow.prompts.agenticrag import AutoPromptGeneratorPrompt
2+
import pandas as pd
3+
from dataflow.utils.registry import OPERATOR_REGISTRY
4+
from dataflow import get_logger
5+
6+
from dataflow.utils.storage import DataFlowStorage
7+
from dataflow.core import OperatorABC
8+
from dataflow.core import LLMServingABC
9+
10+
@OPERATOR_REGISTRY.register()
11+
class AutoPromptGenerator(OperatorABC):
12+
'''
13+
AutoPromptGenerator is a class that generates prompts for given document fragments to generate seed QA pairs.
14+
'''
15+
def __init__(self, llm_serving: LLMServingABC):
16+
self.logger = get_logger()
17+
self.prompts = AutoPromptGeneratorPrompt()
18+
self.llm_serving = llm_serving
19+
20+
@staticmethod
21+
def get_desc(lang: str = "zh"):
22+
if lang == "zh":
23+
return (
24+
"该算子用于为给的的文档片段生成提示词,用于生成种子QA对\n\n"
25+
"输入参数:\n"
26+
"- input_key: 包含文档片段的字段名\n"
27+
"- output_key: 包含提示词的字段名\n"
28+
)
29+
elif lang == "en":
30+
return (
31+
"This operator generates prompts for given document fragments to generate seed QA pairs.\n\n"
32+
"Input Parameters:\n"
33+
"- input_key: Field name containing the content\n"
34+
"- output_key: Field name containing the generated prompt\n"
35+
)
36+
else:
37+
return "AutoPromptGenerator generates prompts for given document fragments to generate seed QA pairs."
38+
39+
def _validate_dataframe(self, dataframe: pd.DataFrame):
40+
required_keys = [self.input_key]
41+
forbidden_keys = [self.output_key]
42+
43+
missing = [k for k in required_keys if k not in dataframe.columns]
44+
conflict = [k for k in forbidden_keys if k in dataframe.columns]
45+
46+
if missing:
47+
raise ValueError(f"Missing required column(s): {missing}")
48+
if conflict:
49+
raise ValueError(f"The following column(s) already exist and would be overwritten: {conflict}")
50+
51+
def _reformat_prompt(self, dataframe):
52+
"""
53+
Reformat the prompts in the dataframe to generate questions.
54+
"""
55+
questions = dataframe[self.input_key].tolist()
56+
inputs = [self.prompts.auto_prompt_generator_prompt(question) for question in questions]
57+
58+
return inputs
59+
60+
def run(
61+
self,
62+
storage: DataFlowStorage,
63+
input_key:str = "text",
64+
output_key:str = "generated_prompt"
65+
):
66+
'''
67+
Runs the answer generation process, reading from the input file and saving results to output.
68+
'''
69+
self.input_key, self.output_key = input_key, output_key
70+
dataframe = storage.read("dataframe")
71+
self._validate_dataframe(dataframe)
72+
formatted_prompts = self._reformat_prompt(dataframe)
73+
answers = self.llm_serving.generate_from_input(user_inputs=formatted_prompts, system_prompt="")
74+
75+
dataframe[self.output_key] = answers
76+
output_file = storage.write(dataframe)
77+
self.logger.info(f"Results saved to {output_file}")
78+
79+
return [output_key]
Lines changed: 92 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1 +1,93 @@
11
import pandas as pd
2+
from dataflow.utils.registry import OPERATOR_REGISTRY
3+
from dataflow import get_logger
4+
5+
from dataflow.utils.storage import DataFlowStorage
6+
from dataflow.core import OperatorABC
7+
from dataflow.core import LLMServingABC
8+
9+
@OPERATOR_REGISTRY.register()
10+
class QAGenerator:
11+
'''
12+
SeedQAGenerator is a class that uses LLMs to generate QA pairs based on seed input.
13+
'''
14+
15+
def __init__(self, llm_serving: LLMServingABC):
16+
self.logger = get_logger()
17+
self.llm_serving = llm_serving
18+
19+
@staticmethod
20+
def get_desc(lang: str = "zh"):
21+
if lang == "zh":
22+
return (
23+
"该算子用于生成对应文档片段的QA对。\n\n"
24+
"输入参数:\n"
25+
"- input_key: 包含文档片段的字段名\n"
26+
"- prompt_key: 包含提示词的字段名\n"
27+
"- output_quesion_key: 包含生成问题的字段名\n"
28+
"- output_answer_key: 包含生成答案的字段名\n"
29+
)
30+
elif lang == "en":
31+
return (
32+
"This operator generates QA pairs for given document fragments.\n\n"
33+
"Input Parameters:\n"
34+
"- input_key: Field name containing the content\n"
35+
"- prompt_key: Field name containing the generated prompt\n"
36+
"- output_quesion_key: Field name containing the generated question\n"
37+
"- output_answer_key: Field name containing the generated answer\n"
38+
)
39+
else:
40+
return "QAGenerator generates QA pairs for given document fragments."
41+
42+
def _validate_dataframe(self, dataframe: pd.DataFrame):
43+
required_keys = [self.input_key]
44+
forbidden_keys = [self.output_question_key, self.output_answer_key]
45+
46+
missing = [k for k in required_keys if k not in dataframe.columns]
47+
conflict = [k for k in forbidden_keys if k in dataframe.columns]
48+
49+
if missing:
50+
raise ValueError(f"Missing required column(s): {missing}")
51+
if conflict:
52+
raise ValueError(f"The following column(s) already exist and would be overwritten: {conflict}")
53+
54+
def _build_prompt(self, df):
55+
prompts = []
56+
for index, row in df.iterrows():
57+
prompts.append(row[self.prompt_key] + "Format:\nQ: ...\nA: ..." + "\nSeed data:\n" + row[self.input_key])
58+
return prompts
59+
60+
def _parse_qa(self, response: str) -> tuple:
61+
lines = response.strip().split('\n')
62+
q = next((line[2:].strip() for line in lines if line.lower().startswith("q:")), "")
63+
a = next((line[2:].strip() for line in lines if line.lower().startswith("a:")), "")
64+
return q, a
65+
66+
def run(
67+
self,
68+
storage: DataFlowStorage,
69+
input_key:str = "text",
70+
output_prompt_key:str = "generated_prompt",
71+
output_quesion_key:str = "generated_question",
72+
output_answer_key:str = "generated_answer"
73+
):
74+
'''
75+
Runs the answer generation process, reading from the input file and saving results to output.
76+
'''
77+
78+
self.input_key, self.prompt_key, self.output_question_key, self.output_answer_key = input_key, output_prompt_key, output_quesion_key, output_answer_key
79+
80+
dataframe = storage.read("dataframe")
81+
self._validate_dataframe(dataframe)
82+
formatted_prompts = self._build_prompt(dataframe)
83+
responses = self.llm_serving.generate_from_input(user_inputs=formatted_prompts, system_prompt="")
84+
85+
questions, answers = zip(*[self._parse_qa(r) for r in responses])
86+
87+
dataframe[self.output_question_key] = questions
88+
dataframe[self.output_answer_key] = answers
89+
90+
output_file = storage.write(dataframe)
91+
self.logger.info(f"Results saved to {output_file}")
92+
93+
return [self.output_question_key, self.output_answer_key]
Lines changed: 154 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1 +1,155 @@
11
from dataflow.prompts.agenticrag import QAScorerPrompt
2+
import pandas as pd
3+
from dataflow.utils.registry import OPERATOR_REGISTRY
4+
from dataflow import get_logger
5+
import re
6+
from dataflow.utils.storage import DataFlowStorage
7+
from dataflow.core import OperatorABC
8+
from dataflow.core import LLMServingABC
9+
10+
@OPERATOR_REGISTRY.register()
11+
class QAScorer(OperatorABC):
12+
'''
13+
Answer Generator is a class that generates answers for given questions.
14+
'''
15+
def __init__(self, llm_serving: LLMServingABC):
16+
self.logger = get_logger()
17+
self.prompts = QAScorerPrompt()
18+
self.llm_serving = llm_serving
19+
20+
@staticmethod
21+
def get_desc(lang: str = "zh"):
22+
if lang == "zh":
23+
return (
24+
"该算子用于为给的的文档片段生成种子QA对打分\n\n"
25+
"输入参数:\n"
26+
"- input_question_key: Field name containing the generated question\n"
27+
"- input_answer_key: Field name containing the generated answer\n"
28+
"- output_question_quality_key: Field name containing the question quality grade\n"
29+
"- output_question_quality_feedback_key: Field name containing the question quality feedback\n"
30+
"- output_answer_alignment_key: Field name containing the answer alignment grade\n"
31+
"- output_answer_alignment_feedback_key: Field name containing the answer alignment feedback\n"
32+
"- output_answer_verifiability_key: Field name containing the answer verifiability grade\n"
33+
"- output_downstream_value_key: Field name containing the downstream value grade\n"
34+
"- output_downstream_value_feedback_key: Field name containing the downstream value feedback\n"
35+
)
36+
elif lang == "en":
37+
return (
38+
"This operator generates prompts for given document fragments to generate seed QA pairs.\n\n"
39+
"Input Parameters:\n"
40+
"- input_question_key: Field name containing the generated question\n"
41+
"- input_answer_key: Field name containing the generated answer\n"
42+
"- output_question_quality_key: Field name containing the question quality grade\n"
43+
"- output_question_quality_feedback_key: Field name containing the question quality feedback\n"
44+
"- output_answer_alignment_key: Field name containing the answer alignment grade\n"
45+
"- output_answer_alignment_feedback_key: Field name containing the answer alignment feedback\n"
46+
"- output_answer_verifiability_key: Field name containing the answer verifiability grade\n"
47+
"- output_downstream_value_key: Field name containing the downstream value grade\n"
48+
"- output_downstream_value_feedback_key: Field name containing the downstream value feedback\n"
49+
)
50+
else:
51+
return "QAScorer scores QA pairs for given document fragments."
52+
53+
def _validate_dataframe(self, dataframe: pd.DataFrame):
54+
required_keys = [self.input_question_key, self.input_answer_key]
55+
forbidden_keys = [self.output_question_quality_key, self.output_question_quality_feedback_key, self.output_answer_alignment_key, self.output_answer_alignment_feedback_key, self.output_answer_verifiability_key, self.output_answer_verifiability_feedback_key, self.output_downstream_value_key, self.output_downstream_value_feedback_key]
56+
57+
missing = [k for k in required_keys if k not in dataframe.columns]
58+
conflict = [k for k in forbidden_keys if k in dataframe.columns]
59+
60+
if missing:
61+
raise ValueError(f"Missing required column(s): {missing}")
62+
if conflict:
63+
raise ValueError(f"The following column(s) already exist and would be overwritten: {conflict}")
64+
65+
def _build_prompts(self, dataframe):
66+
"""
67+
Reformat the prompts in the dataframe to generate questions.
68+
"""
69+
question_quality_inputs = []
70+
question_quality_prompt = self.prompts.question_quality_prompt()
71+
answer_alignment_inputs = []
72+
answer_alignment_prompt = self.prompts.answer_alignment_prompt()
73+
answer_verifiability_inputs = []
74+
answer_verifiability_prompt = self.prompts.answer_verifiability_prompt()
75+
downstream_value_inputs = []
76+
downstream_value_prompt = self.prompts.downstream_value_prompt()
77+
78+
for index, row in dataframe.iterrows():
79+
question_quality_content = question_quality_prompt + "Question: " + row[self.input_question_key] + "\n" + "Answer: " + row[self.input_answer_key]
80+
question_quality_inputs.append(question_quality_content)
81+
answer_alignment_content = answer_alignment_prompt + "Question: " + row[self.input_question_key] + "\n" + "Answer: " + row[self.input_answer_key]
82+
answer_alignment_inputs.append(answer_alignment_content)
83+
answer_verifiability_content = answer_verifiability_prompt + "Question: " + row[self.input_question_key] + "\n" + "Answer: " + row[self.input_answer_key]
84+
answer_verifiability_inputs.append(answer_verifiability_content)
85+
downstream_value_content = downstream_value_prompt + "Question: " + row[self.input_question_key] + "\n" + "Answer: " + row[self.input_answer_key]
86+
downstream_value_inputs.append(downstream_value_content)
87+
88+
return question_quality_inputs, answer_alignment_inputs, answer_verifiability_inputs, downstream_value_inputs
89+
90+
def _parse_grade_and_feedback(self, response: str) -> tuple:
91+
grading_match = re.search(r"\*\*Grading\*\*:\s*(\d+)", response)
92+
feedback_match = re.search(r"\*\*Feedback\*\*:\s*(.+)", response, re.DOTALL)
93+
grading = float(grading_match.group(1)) if grading_match else 0
94+
feedback = feedback_match.group(1).strip() if feedback_match else ''
95+
96+
return grading, feedback
97+
98+
def run(
99+
self,
100+
storage: DataFlowStorage,
101+
input_question_key: str = "generated_question",
102+
input_answer_key: str = "generated_answer",
103+
output_question_quality_key: str = "question_quality_grades",
104+
output_question_quality_feedback_key: str = "question_quality_feedbacks",
105+
output_answer_alignment_key: str = "answer_alignment_grades",
106+
output_answer_alignment_feedback_key: str = "answer_alignment_feedbacks",
107+
output_answer_verifiability_key: str = "answer_verifiability_grades",
108+
output_answer_verifiability_feedback_key: str = "answer_verifiability_feedbacks",
109+
output_downstream_value_key: str = "downstream_value_grades",
110+
output_downstream_value_feedback_key: str = "downstream_value_feedbacks"
111+
):
112+
self.input_question_key, self.input_answer_key, self.output_question_quality_key, self.output_question_quality_feedback_key, self.output_answer_alignment_key, self.output_answer_alignment_feedback_key, self.output_answer_verifiability_key, self.output_answer_verifiability_feedback_key, self.output_downstream_value_key, self.output_downstream_value_feedback_key = input_question_key, input_answer_key, output_question_quality_key, output_question_quality_feedback_key, output_answer_alignment_key, output_answer_alignment_feedback_key, output_answer_verifiability_key, output_answer_verifiability_feedback_key, output_downstream_value_key, output_downstream_value_feedback_key
113+
114+
dataframe = storage.read("dataframe")
115+
self._validate_dataframe(dataframe)
116+
117+
# 构建prompt
118+
q_inputs, a_inputs, v_inputs, d_inputs = self._build_prompts(dataframe)
119+
120+
# 生成四类分数和反馈
121+
self.logger.info("Scoring question quality...")
122+
q_scores = self.llm_serving.generate_from_input(user_inputs=q_inputs, system_prompt="")
123+
q_grades, q_feedbacks = zip(*[self._parse_grade_and_feedback(r) for r in q_scores])
124+
125+
self.logger.info("Scoring answer alignment...")
126+
a_scores = self.llm_serving.generate_from_input(user_inputs=a_inputs, system_prompt="")
127+
a_grades, a_feedbacks = zip(*[self._parse_grade_and_feedback(r) for r in a_scores])
128+
129+
self.logger.info("Scoring answer verifiability...")
130+
v_scores = self.llm_serving.generate_from_input(user_inputs=v_inputs, system_prompt="")
131+
v_grades, v_feedbacks = zip(*[self._parse_grade_and_feedback(r) for r in v_scores])
132+
133+
self.logger.info("Scoring downstream value...")
134+
d_scores = self.llm_serving.generate_from_input(user_inputs=d_inputs, system_prompt="")
135+
d_grades, d_feedbacks = zip(*[self._parse_grade_and_feedback(r) for r in d_scores])
136+
137+
# 写回结果
138+
dataframe[self.output_question_quality_key] = q_grades
139+
dataframe[self.output_question_quality_feedback_key] = q_feedbacks
140+
dataframe[self.output_answer_alignment_key] = a_grades
141+
dataframe[self.output_answer_alignment_feedback_key] = a_feedbacks
142+
dataframe[self.output_answer_verifiability_key] = v_grades
143+
dataframe[self.output_answer_verifiability_feedback_key] = v_feedbacks
144+
dataframe[self.output_downstream_value_key] = d_grades
145+
dataframe[self.output_downstream_value_feedback_key] = d_feedbacks
146+
147+
output_file = storage.write(dataframe)
148+
self.logger.info(f"Results saved to {output_file}")
149+
150+
return [
151+
output_question_quality_key, output_question_quality_feedback_key,
152+
output_answer_alignment_key, output_answer_alignment_feedback_key,
153+
output_answer_verifiability_key, output_answer_verifiability_feedback_key,
154+
output_downstream_value_key, output_downstream_value_feedback_key
155+
]

0 commit comments

Comments
 (0)