From 262e0d6df0ac827806ddeca6c10a33cb4d57b4d7 Mon Sep 17 00:00:00 2001 From: slin000111 Date: Tue, 12 Dec 2023 19:54:02 +0800 Subject: [PATCH] fix embedding and inference device in faq question answering pipeline --- modelscope/models/nlp/structbert/faq_question_answering.py | 2 ++ modelscope/pipelines/nlp/faq_question_answering_pipeline.py | 3 +++ 2 files changed, 5 insertions(+) diff --git a/modelscope/models/nlp/structbert/faq_question_answering.py b/modelscope/models/nlp/structbert/faq_question_answering.py index bc22ab61..6c05bcff 100644 --- a/modelscope/models/nlp/structbert/faq_question_answering.py +++ b/modelscope/models/nlp/structbert/faq_question_answering.py @@ -375,6 +375,8 @@ class ProtoNet(nn.Module): input_ids = torch.IntTensor(input_ids) if not isinstance(input_mask, Tensor): input_mask = torch.IntTensor(input_mask) + input_ids = input_ids.to(self.bert.device) + input_mask = input_mask.to(self.bert.device) rst = self.bert(input_ids, input_mask) last_hidden_states = rst.last_hidden_state if len(input_mask.shape) == 2: diff --git a/modelscope/pipelines/nlp/faq_question_answering_pipeline.py b/modelscope/pipelines/nlp/faq_question_answering_pipeline.py index 0b2ba199..3205f8b5 100644 --- a/modelscope/pipelines/nlp/faq_question_answering_pipeline.py +++ b/modelscope/pipelines/nlp/faq_question_answering_pipeline.py @@ -50,6 +50,9 @@ class FaqQuestionAnsweringPipeline(Pipeline): return pipeline_parameters, pipeline_parameters, pipeline_parameters def get_sentence_embedding(self, inputs, max_len=None): + if (self.model or (self.has_multiple_models and self.models[0])): + if not self._model_prepare: + self.prepare_model() inputs = self.preprocessor.batch_encode(inputs, max_length=max_len) sentence_vecs = self.model.forward_sentence_embedding(inputs) sentence_vecs = sentence_vecs.detach().tolist()