Przeglądaj źródła

Merge branch 'dyx' of ChenJiaWei/knowledge-backend into master

DingYunXiang 3 lat temu
rodzic
commit
1c4a6b74a1

+ 22 - 0
src/main/java/com/seeckg/knowledgegraph_backend/controller/QuestionAnsweringController.java

@@ -0,0 +1,22 @@
+package com.seeckg.knowledgegraph_backend.controller;
+
+import com.seeckg.knowledgegraph_backend.pojo.Knowledge;
+import com.seeckg.knowledgegraph_backend.pojo.Response;
+import com.seeckg.knowledgegraph_backend.service.QuestionAnsweringService;
+import org.springframework.beans.factory.annotation.Autowired;
+import org.springframework.web.bind.annotation.GetMapping;
+import org.springframework.web.bind.annotation.PathVariable;
+import org.springframework.web.bind.annotation.RequestMapping;
+import org.springframework.web.bind.annotation.RestController;
+
+@RestController
+@RequestMapping("/api/question-answering")
+public class QuestionAnsweringController {
+    @Autowired
+    QuestionAnsweringService questionAnsweringService;
+
+    @GetMapping("/{type}/{question}")
+    public Response findByName(@PathVariable("type") int type, @PathVariable("question") String question) {
+        return questionAnsweringService.getAnswer(type, question);
+    }
+}

+ 12 - 0
src/main/java/com/seeckg/knowledgegraph_backend/enums/questionAnsweringType.java

@@ -0,0 +1,12 @@
+package com.seeckg.knowledgegraph_backend.enums;
+
+public enum questionAnsweringType {
+    /**
+     * 离线问答
+     */
+    OFFLINE,
+    /**
+     * 在线问答
+     */
+    ONLINE
+}

+ 10 - 0
src/main/java/com/seeckg/knowledgegraph_backend/remote/GPTRemoteService.java

@@ -0,0 +1,10 @@
+package com.seeckg.knowledgegraph_backend.remote;
+
+import org.springframework.stereotype.Service;
+
+@Service
+public class GPTRemoteService {
+    public String getGPTAnswer(String question) {
+        return "";
+    }
+}

+ 4 - 0
src/main/java/com/seeckg/knowledgegraph_backend/service/KnowledgeService.java

@@ -128,6 +128,10 @@ public class KnowledgeService {
         return knowledgeDao.findByName(name);
     }
 
+    public List<Knowledge> getAllKnowledge() {
+        return knowledgeDao.findAll();
+    }
+
     /**
      *
      * 添加关系

+ 90 - 0
src/main/java/com/seeckg/knowledgegraph_backend/service/QuestionAnsweringService.java

@@ -0,0 +1,90 @@
+package com.seeckg.knowledgegraph_backend.service;
+
+
+import com.seeckg.knowledgegraph_backend.enums.ReturnCode;
+import com.seeckg.knowledgegraph_backend.enums.questionAnsweringType;
+import com.seeckg.knowledgegraph_backend.pojo.Knowledge;
+import com.seeckg.knowledgegraph_backend.pojo.Response;
+import com.seeckg.knowledgegraph_backend.remote.GPTRemoteService;
+import lombok.extern.slf4j.Slf4j;
+import org.springframework.stereotype.Service;
+
+import javax.annotation.Resource;
+import java.io.BufferedReader;
+import java.io.IOException;
+import java.io.InputStreamReader;
+import java.util.ArrayList;
+import java.util.List;
+
+@Service
+@Slf4j
+public class QuestionAnsweringService {
+    @Resource
+    GPTRemoteService gptRemoteService;
+    KnowledgeService knowledgeService;
+
+
+    /**
+     * 选择离线问答或在线问答
+     * 返回的数据格式是String
+     * @return 包含上述数据的Response
+     */
+    public Response getAnswer(int type, String question) {
+        if (type == questionAnsweringType.OFFLINE.ordinal()) {
+            return getOfflineAnswer(question);
+        }
+        else if (type == questionAnsweringType.ONLINE.ordinal()) {
+            return getOnlineAnswer(question);
+        }
+        return Response.buildFailed(ReturnCode.PARAM_ERROR.getCode(), "问答类型不存在!");
+    }
+
+
+    /**
+     * 获取离线问答结果
+     * 返回的数据格式是String
+     * @return 包含上述数据的Response
+     */
+    public Response getOfflineAnswer(String question) {
+        String answer = "";
+
+        List<Knowledge> allKnowledge = knowledgeService.getAllKnowledge();
+        List<String> allKnowledgeText = new ArrayList<>();
+        for (Knowledge knowledge : allKnowledge) {
+            String context = knowledge.getContext().replace('\n', ' ');
+            allKnowledgeText.add(context);
+        }
+        String contexts = String.join(" ", allKnowledgeText);
+
+        try {
+            String[] args1 = new String[] { "python", "..\\util\\OfflineQuestionAnsweringServive.py", question, contexts};
+            Process proc = Runtime.getRuntime().exec(args1);
+
+            BufferedReader in = new BufferedReader(new InputStreamReader(proc.getInputStream()));
+            String line = null;
+            while ((line = in.readLine()) != null) {
+                answer = line;
+                // System.out.println(line);
+            }
+            in.close();
+            proc.waitFor();
+        } catch (IOException | InterruptedException e) {
+            e.printStackTrace();
+            return Response.buildFailed(ReturnCode.SERVER_ERROR.getCode(), "答案生成出错!");
+        }
+
+        return Response.buildSuccess(answer);
+    }
+
+
+    /**
+     * 获取在线问答结果
+     * 返回的数据格式是String
+     * @return 包含上述数据的Response
+     */
+    public Response getOnlineAnswer(String question) {
+        String answer = gptRemoteService.getGPTAnswer(question);
+        return Response.buildSuccess(answer);
+    }
+
+}

+ 27 - 0
src/main/java/com/seeckg/knowledgegraph_backend/util/OfflineQuestionAnsweringServive.py

@@ -0,0 +1,27 @@
+import sys
+from transformers import AutoModelForQuestionAnswering, AutoTokenizer, pipeline
+
+
+def getOfflineAnswer(question, contexts):
+    """
+    model = AutoModelForQuestionAnswering.from_pretrained('uer/roberta-base-chinese-extractive-qa')
+    tokenizer = AutoTokenizer.from_pretrained('uer/roberta-base-chinese-extractive-qa')
+    """
+    model = AutoModelForQuestionAnswering.from_pretrained('luhua/chinese_pretrain_mrc_roberta_wwm_ext_large')
+    tokenizer = AutoTokenizer.from_pretrained('luhua/chinese_pretrain_mrc_roberta_wwm_ext_large')
+
+    qa_model = pipeline("question-answering", model=model, tokenizer=tokenizer)
+    answer = qa_model(question=question, context=contexts, max_answer_len=128, doc_stride=256)
+    print(answer)
+    return dict(answer).get('answer', '')
+
+
+if __name__ == '__main__':
+    question = sys.argv[1]
+    contexts = []
+    for i in range(2, len(sys.argv)):
+        contexts.append(sys.argv[i])
+    contexts = ' '.join(contexts)
+    # question = "软件工程的定义?"
+    # contexts = "软件工程是一门研究用工程化方法构建和维护有效、实用和高质量的软件的学科。它涉及程序设计语言、数据库、软件开发工具、系统平台、标准、设计件有电子邮件、嵌入式系统、人机界面、办公套件、操作系统、编译器、数据库、游戏等。同时,各个行业几乎都有计算机软件的应用,如工业、农业、银行、航空、政府部门等。这些应用促进了经济和社会的发展,也提高了工作效率和生活效率。"
+    print(getOfflineAnswer(question, contexts))