INNER CODE UNIT · Python

predict_list

JackHCC/Chinese-Text-Classification-PyTorch · predict.py:67

    def predict_list(self, querys):
        # 返回预测的索引
        data = self.build_predict_text(querys)
        with torch.no_grad():
            outputs = self.model(data)
            num = torch.argmax(outputs, dim=1)
            pred = [key[index] for index in list(np.array(num))]
        return pred


if __name__ == "__main__":
    pred = Predict('TextCNN')
    # 预测一条
    query = "学费太贵怎么办?"
    print(pred.predict(query))
    # 预测一个列表
    querys = ["学费太贵怎么办?", "金融怎么样"]
    print(pred.predict_list(querys))

View source record →

📰 Research Paper
Loading…
⏳ Fetching content…