form.py 1.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051
  1. # -*- coding: utf-8 -*-
  2. """``FormPredictor`` — 表格预测。
  3. Phase 5 从 ``interface/predictor.py``(约 1578-1612 行)迁出。
  4. 原 ``from common.Utils import *`` / ``from interface.modelFactory import *``
  5. 已替换为显式 import;``os.path.dirname(__file__)`` 路径引用替换为
  6. ``predictors._common.INTERFACE_DIR``。
  7. """
  8. from __future__ import absolute_import
  9. import os
  10. from BiddingKG.dl.common.Utils import encodeInput, encodeInput_form
  11. from BiddingKG.dl.interface.modelFactory import Model_form_item, Model_form_context
  12. from BiddingKG.dl.model_runtime.embed import getLazyLoad
  13. from BiddingKG.dl.predictors._common import INTERFACE_DIR
  14. __all__ = ["FormPredictor"]
  15. #表格预测
  16. class FormPredictor():
  17. def __init__(self,lazyLoad=getLazyLoad(),config=None):
  18. self.model_file_line = os.path.join(INTERFACE_DIR, "..", "form", "model", "model_form.model_line.hdf5")
  19. self.model_file_item = os.path.join(INTERFACE_DIR, "..", "form", "model", "model_form.model_item.hdf5")
  20. self.model_form_item = Model_form_item(config=config)
  21. self.model_dict = {"line":[None,self.model_file_line]}
  22. self.model_form_context = Model_form_context(config=config)
  23. def getModel(self,type):
  24. if type=="item":
  25. return self.model_form_item
  26. elif type=="context":
  27. return self.model_form_context
  28. else:
  29. return self.getModel(type)
  30. def encode(self,data,**kwargs):
  31. return encodeInput([data], word_len=50, word_flag=True,userFool=False)[0]
  32. return encodeInput_form(data)
  33. def predict(self,form_datas,type):
  34. if type=="item":
  35. return self.model_form_item.predict(form_datas)
  36. elif type=="context":
  37. return self.model_form_context.predict(form_datas)
  38. else:
  39. return self.getModel(type).predict(form_datas)