11from 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