Java中Python Transformers包替代方案及HuggingFace QA模型加载方法
在Java中加载问答预训练模型并搭建API
可用的Java依赖
你可以使用Hugging Face Transformers Java库,这是官方提供的Java版Transformers实现,支持加载Hugging Face上的预训练模型,包括你提到的deepset/roberta-base-squad2问答模型。
步骤1:引入Maven依赖
在你的pom.xml中添加以下依赖:
<dependency> <groupId>ai.djl.huggingface</groupId> <artifactId>transformers</artifactId> <version>0.25.0</version> </dependency> <dependency> <groupId>ai.djl.huggingface</groupId> <artifactId>tokenizers</artifactId> <version>0.25.0</version> </dependency>
(注:版本号可根据最新稳定版调整)
步骤2:加载模型并执行问答预测
以下是加载deepset/roberta-base-squad2模型、接收question和context参数并返回答案的核心代码:
import ai.djl.huggingface.tokenizers.HuggingFaceTokenizer; import ai.djl.modality.nlp.qa.QAInput; import ai.djl.repository.zoo.Criteria; import ai.djl.training.util.ProgressBar; import ai.djl.transformers.qa.QAProcessor; import ai.djl.transformers.qa.QAResult; import ai.djl.transformers.qa.QuestionAnsweringModel; public class QAModelDemo { public static void main(String[] args) throws Exception { // 定义模型名称 String modelName = "deepset/roberta-base-squad2"; // 构建加载模型的Criteria Criteria<QAInput, QAResult> criteria = Criteria.builder() .setTypes(QAInput.class, QAResult.class) .optModelUrls(modelName) .optTranslator(new QAProcessor(modelName)) .optProgress(new ProgressBar()) .build(); // 加载模型 try (QuestionAnsweringModel model = criteria.loadModel().newModel()) { // 模拟用户传入的参数 String question = "Why is model conversion important?"; String context = "The option to convert models between FARM and transformers"; // 构建输入 QAInput input = new QAInput(question, context); // 获取预测结果 QAResult result = model.predict(input); // 输出结果 System.out.println("答案: " + result.getAnswer()); System.out.println("置信度: " + result.getScore()); } } }
步骤3:搭建接收参数的API(基于Spring Boot)
如果要搭建一个接收HTTP请求的API,可以结合Spring Boot实现:
- 先添加Spring Boot Web依赖到
pom.xml:
<dependency> <groupId>org.springframework.boot</groupId> <artifactId>spring-boot-starter-web</artifactId> <version>3.2.0</version> </dependency>
- 编写API接口:
import ai.djl.repository.zoo.Criteria; import ai.djl.training.util.ProgressBar; import ai.djl.transformers.qa.QAProcessor; import ai.djl.transformers.qa.QAResult; import ai.djl.transformers.qa.QuestionAnsweringModel; import ai.djl.modality.nlp.qa.QAInput; import org.springframework.web.bind.annotation.PostMapping; import org.springframework.web.bind.annotation.RequestBody; import org.springframework.web.bind.annotation.RestController; import javax.annotation.PostConstruct; @RestController public class QAApiController { private QuestionAnsweringModel qaModel; private final String modelName = "deepset/roberta-base-squad2"; // 初始化时加载模型 @PostConstruct public void initModel() throws Exception { Criteria<QAInput, QAResult> criteria = Criteria.builder() .setTypes(QAInput.class, QAResult.class) .optModelUrls(modelName) .optTranslator(new QAProcessor(modelName)) .optProgress(new ProgressBar()) .build(); qaModel = criteria.loadModel().newModel(); } // 定义接收参数的接口 @PostMapping("/qa") public QAResult getAnswer(@RequestBody QAInput input) throws Exception { return qaModel.predict(input); } } // 用于接收请求参数的实体类 class QAInput { private String question; private String context; // Getter和Setter public String getQuestion() { return question; } public void setQuestion(String question) { this.question = question; } public String getContext() { return context; } public void setContext(String context) { this.context = context; } }
启动Spring Boot应用后,就可以通过POST请求/qa接口,传入包含question和context的JSON参数,获取问答结果。
内容的提问来源于stack exchange,提问作者Ahmad Mujtaba
相关产品推荐
相关产品推荐

