| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051 |
- # -*- coding: utf-8 -*-
- """``FormPredictor`` — 表格预测。
- Phase 5 从 ``interface/predictor.py``(约 1578-1612 行)迁出。
- 原 ``from common.Utils import *`` / ``from interface.modelFactory import *``
- 已替换为显式 import;``os.path.dirname(__file__)`` 路径引用替换为
- ``predictors._common.INTERFACE_DIR``。
- """
- from __future__ import absolute_import
- import os
- from BiddingKG.dl.common.Utils import encodeInput, encodeInput_form
- from BiddingKG.dl.interface.modelFactory import Model_form_item, Model_form_context
- from BiddingKG.dl.model_runtime.embed import getLazyLoad
- from BiddingKG.dl.predictors._common import INTERFACE_DIR
- __all__ = ["FormPredictor"]
- #表格预测
- class FormPredictor():
-
- def __init__(self,lazyLoad=getLazyLoad(),config=None):
- self.model_file_line = os.path.join(INTERFACE_DIR, "..", "form", "model", "model_form.model_line.hdf5")
- self.model_file_item = os.path.join(INTERFACE_DIR, "..", "form", "model", "model_form.model_item.hdf5")
- self.model_form_item = Model_form_item(config=config)
- self.model_dict = {"line":[None,self.model_file_line]}
- self.model_form_context = Model_form_context(config=config)
-
- def getModel(self,type):
- if type=="item":
- return self.model_form_item
- elif type=="context":
- return self.model_form_context
- else:
- return self.getModel(type)
- def encode(self,data,**kwargs):
- return encodeInput([data], word_len=50, word_flag=True,userFool=False)[0]
- return encodeInput_form(data)
-
- def predict(self,form_datas,type):
- if type=="item":
- return self.model_form_item.predict(form_datas)
- elif type=="context":
- return self.model_form_context.predict(form_datas)
- else:
- return self.getModel(type).predict(form_datas)
|