From b41726f04fa610d94dd92245b996c5ee7b069ccb Mon Sep 17 00:00:00 2001 From: hp <794731517@qq.com> Date: Sun, 21 Aug 2022 09:54:28 +0800 Subject: [PATCH 1/2] =?UTF-8?q?=E4=BF=AE=E6=94=B9=E6=8E=A8=E8=8D=90?= =?UTF-8?q?=E9=80=BB=E8=BE=91=EF=BC=8C=E5=A2=9E=E5=8A=A0=E5=AE=89=E8=A3=85?= =?UTF-8?q?=E9=98=88=E5=80=BC=E5=A2=9E=E5=8A=A0=E6=8E=A8=E8=8D=90=E7=AD=94?= =?UTF-8?q?=E6=A1=88?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/model/q_a.py | 22 ++++++++++++++-------- src/qaRobot/qa_api/qa/q_a.py | 15 +++++++++------ src/qaRobot/qa_api/views.py | 12 ++++++++---- 3 files changed, 31 insertions(+), 18 deletions(-) diff --git a/src/model/q_a.py b/src/model/q_a.py index 3a97829..e9b8968 100644 --- a/src/model/q_a.py +++ b/src/model/q_a.py @@ -10,16 +10,20 @@ context.set_context(mode=context.GRAPH_MODE, device_target="CPU") def compute_similarity(input_encode): - max_similarity = 0 with open("../data/resource_sentence_encode.json") as f: resource_sentences_encode = json.load(f) + question_similarity = {} for k, v in resource_sentences_encode.items(): similarity = cosine_similarity( [input_encode[1][0].asnumpy()], [np.asarray(v)]) - if similarity > max_similarity: - max_similarity = similarity - match_key = k - return match_key + question_similarity[k] = similarity + sorted_similarity = sorted(question_similarity.items(), key=lambda x: x[1], reverse=True) + if not sorted_similarity or len(sorted_similarity) == 0: + return None + elif sorted_similarity[0][1] >= 0.7: + return [sorted_similarity[0][0]] + elif len(sorted_similarity) > 1 and sorted_similarity[0][1] < 0.7: + return [sorted_similarity[0][0], sorted_similarity[1][0]] def encode_sentence(input_sentence): @@ -38,7 +42,9 @@ def load_q_a_data(file_path): if __name__ == '__main__': input_sentence = sys.argv[1] input_encode = encode_sentence(input_sentence) - match_key = compute_similarity(input_encode) + match_keys = compute_similarity(input_encode) q_a_data = load_q_a_data("../data/q_a.json") - if match_key in q_a_data.keys(): - print(q_a_data[match_key]) + if match_keys: + for k in match_keys: + if k in q_a_data.keys(): + print(q_a_data[k]) diff --git a/src/qaRobot/qa_api/qa/q_a.py b/src/qaRobot/qa_api/qa/q_a.py index 2f55a9a..9044cd4 100644 --- a/src/qaRobot/qa_api/qa/q_a.py +++ b/src/qaRobot/qa_api/qa/q_a.py @@ -12,17 +12,20 @@ module_dir = os.path.dirname(__file__) def get_match_question(input_encode): - max_similarity = 0 - print(os.path.join(os.path.dirname(os.path.dirname(__file__)))) with open(os.path.join(module_dir, '../data/resource_sentence_encode.json')) as f: resource_sentences_encode = json.load(f) + question_similarity = {} for k, v in resource_sentences_encode.items(): similarity = cosine_similarity( [input_encode[1][0].asnumpy()], [np.asarray(v)]) - if similarity > max_similarity: - max_similarity = similarity - match_key = k - return match_key + question_similarity[k] = similarity + sorted_similarity = sorted(question_similarity.items(), key=lambda x: x[1], reverse=True) + if not sorted_similarity or len(sorted_similarity) == 0: + return None + elif sorted_similarity[0][1] >= 0.7: + return [sorted_similarity[0][0]] + elif len(sorted_similarity) > 1 and sorted_similarity[0][1] < 0.7: + return [sorted_similarity[0][0], sorted_similarity[1][0]] def encode_sentence(input_sentence): diff --git a/src/qaRobot/qa_api/views.py b/src/qaRobot/qa_api/views.py index 37b39d4..79cf37e 100644 --- a/src/qaRobot/qa_api/views.py +++ b/src/qaRobot/qa_api/views.py @@ -13,12 +13,16 @@ module_dir = os.path.dirname(__file__) def get_answer(request): question = request.query_params.get('q', None) input_encode = encode_sentence(question) - match_key = get_match_question(input_encode) + match_keys = get_match_question(input_encode) file_path = os.path.join(module_dir, 'data/q_a.json') q_a_data = {} if not q_a_data: q_a_data = load_q_a_data(file_path) - r = {"question": question} - if match_key in q_a_data.keys(): - r["answer"] = q_a_data[match_key] + r = [] + if match_keys: + for k in match_keys: + q_a = {"question": question} + if k in q_a_data.keys(): + q_a["answer"] = q_a_data[k] + r.append(q_a) return HttpResponse(json.dumps(r, ensure_ascii=False), content_type='application/json') -- Gitee From 2e3b66186af28f6d9e3f47cc69d6493e9eb87288 Mon Sep 17 00:00:00 2001 From: hp <794731517@qq.com> Date: Sun, 21 Aug 2022 10:27:46 +0800 Subject: [PATCH 2/2] =?UTF-8?q?=E4=BF=AE=E6=94=B9=E6=8E=A8=E8=8D=90?= =?UTF-8?q?=E9=80=BB=E8=BE=91=EF=BC=8C=E5=A2=9E=E5=8A=A0=E5=AE=89=E8=A3=85?= =?UTF-8?q?=E9=98=88=E5=80=BC=E5=A2=9E=E5=8A=A0=E6=8E=A8=E8=8D=90=E7=AD=94?= =?UTF-8?q?=E6=A1=88?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- src/qaRobot/qa_api/views.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/qaRobot/qa_api/views.py b/src/qaRobot/qa_api/views.py index 79cf37e..8958aad 100644 --- a/src/qaRobot/qa_api/views.py +++ b/src/qaRobot/qa_api/views.py @@ -21,7 +21,7 @@ def get_answer(request): r = [] if match_keys: for k in match_keys: - q_a = {"question": question} + q_a = {"question": k} if k in q_a_data.keys(): q_a["answer"] = q_a_data[k] r.append(q_a) -- Gitee