diff --git a/src/model/q_a.py b/src/model/q_a.py index 3a97829c2561b65ab4061640b5f61506e1619f2d..e9b8968641654e7f5e8fd87965e82007f8761ac0 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 2f55a9ab122daccdbe08bc52f1b6efcf023b0982..9044cd4741ee9cba49d03c0f03ae585357ee859c 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 37b39d4d89064b2787c181c74f48360d221e2cbb..8958aade0af7887614c69505bae9a16bc3ea1b4d 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": k} + 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')