paieas.py 6.1 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154
  1. # -*- coding: utf-8 -*-
  2. """PAI-EAS 外部推理服务客户端。
  3. 按 ARCHITECTURE.md Phase 3 拆分建议,从 ``common/Utils.py`` 迁出。
  4. 类型:INFRA(外部服务适配)。
  5. 原位置:``common/Utils.py`` 中以下内容:
  6. - ``USE_PAI_EAS`` / ``API_URL`` / ``USE_API`` 全局变量(Phase 1 已改为从 infra/config 读取)
  7. - PAI-EAS 各模型的 authorization / url 常量
  8. - ``limitRun`` — TF session 分批推理
  9. - ``get_values`` — PAI-EAS 响应解析
  10. - ``vpc_requests`` — PAI-EAS VPC 请求
  11. ``common/Utils.py`` 仍 re-export 以上全部名称,老 import 不受影响。
  12. 注意:
  13. ``tf_predict_pb2`` 原代码中被注释掉(``# from pai_tf_predict_proto import tf_predict_pb2``),
  14. 本文件保持相同行为:尝试 import,失败则设为 None,``get_values`` / ``vpc_requests``
  15. 在调用时会因 ``tf_predict_pb2 is None`` 而报错,与原实现行为一致。
  16. """
  17. from __future__ import absolute_import
  18. import numpy as np
  19. import requests
  20. # PAI-EAS / 内部 API 配置走 infra/config(Phase 1)
  21. try:
  22. from BiddingKG.dl.infra import config as _infra_config
  23. _use_pai, _api_url, _use_api = _infra_config.pai_eas_settings()
  24. USE_PAI_EAS = _use_pai
  25. API_URL = _api_url or "http://127.0.0.1:888"
  26. USE_API = _use_api
  27. except Exception:
  28. # infra 不可用时回退默认值,保证其他功能不挂
  29. USE_PAI_EAS = False
  30. API_URL = "http://127.0.0.1:888"
  31. USE_API = False
  32. # pai_tf_predict_proto 可能在运行环境不可用,保持与原 Utils.py 一致的延迟/容错行为
  33. try:
  34. from pai_tf_predict_proto import tf_predict_pb2
  35. except ImportError:
  36. tf_predict_pb2 = None
  37. __all__ = [
  38. "USE_PAI_EAS",
  39. "API_URL",
  40. "USE_API",
  41. "tf_predict_pb2",
  42. "selffool_authorization",
  43. "selffool_url",
  44. "selffool_seg_authorization",
  45. "selffool_seg_url",
  46. "codename_authorization",
  47. "codename_url",
  48. "form_item_authorization",
  49. "form_item_url",
  50. "person_authorization",
  51. "person_url",
  52. "role_authorization",
  53. "role_url",
  54. "money_authorization",
  55. "money_url",
  56. "codeclasses_authorization",
  57. "codeclasses_url",
  58. "limitRun",
  59. "get_values",
  60. "vpc_requests",
  61. ]
  62. selffool_authorization = "NjlhMWFjMjVmNWYyNzI0MjY1OGQ1M2Y0ZmY4ZGY0Mzg3Yjc2MTVjYg=="
  63. selffool_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/selffool_gpu"
  64. selffool_seg_authorization = "OWUwM2Q0ZmE3YjYxNzU4YzFiMjliNGVkMTA3MzJkNjQ2MzJiYzBhZg=="
  65. selffool_seg_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/selffool_seg_gpu"
  66. codename_authorization = "Y2M5MDUxMzU1MTU4OGM3ZDk2ZmEzYjkxYmYyYzJiZmUyYTgwYTg5NA=="
  67. codename_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/codename_gpu"
  68. form_item_authorization = "ODdkZWY1YWY0NmNhNjU2OTI2NWY4YmUyM2ZlMDg1NTZjOWRkYTVjMw=="
  69. form_item_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/form"
  70. person_authorization = "N2I2MDU2N2Q2MGQ0ZWZlZGM3NDkyNTA1Nzc4YmM5OTlhY2MxZGU1Mw=="
  71. person_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/person"
  72. role_authorization = "OWM1ZDg5ZDEwYTEwYWI4OGNjYmRlMmQ1NzYwNWNlZGZkZmRmMjE4OQ=="
  73. role_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/role"
  74. money_authorization = "MDQyNjc2ZDczYjBhYmM4Yzc4ZGI4YjRmMjc3NGI5NTdlNzJiY2IwZA=="
  75. money_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/money"
  76. codeclasses_authorization = "MmUyNWIxZjQ2NjAzMWJlMGIzYzkxMjMzNWY5OWI3NzJlMWQ1ZjY4Yw=="
  77. codeclasses_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/codeclasses"
  78. def limitRun(sess, list_output, feed_dict, MAX_BATCH=1024):
  79. len_sample = 0
  80. if len(feed_dict.keys()) > 0:
  81. len_sample = len(feed_dict[list(feed_dict.keys())[0]])
  82. if len_sample > MAX_BATCH:
  83. list_result = [[] for _ in range(len(list_output))]
  84. _begin = 0
  85. while _begin < len_sample:
  86. new_dict = dict()
  87. for _key in feed_dict.keys():
  88. if isinstance(feed_dict[_key], (float, int, np.int32, np.float_, np.float16, np.float32, np.float64)):
  89. new_dict[_key] = feed_dict[_key]
  90. else:
  91. new_dict[_key] = feed_dict[_key][_begin:_begin + MAX_BATCH]
  92. _output = sess.run(list_output, feed_dict=new_dict)
  93. for _index in range(len(list_output)):
  94. list_result[_index].extend(_output[_index])
  95. _begin += MAX_BATCH
  96. else:
  97. list_result = sess.run(list_output, feed_dict=feed_dict)
  98. return list_result
  99. def get_values(response, output_name):
  100. """
  101. Get the value of a specified output tensor
  102. :param output_name: name of the output tensor
  103. :return: the content of the output tensor
  104. """
  105. output = response.outputs[output_name]
  106. if output.dtype == tf_predict_pb2.DT_FLOAT:
  107. _value = output.float_val
  108. elif output.dtype == tf_predict_pb2.DT_INT8 or output.dtype == tf_predict_pb2.DT_INT16 or \
  109. output.dtype == tf_predict_pb2.DT_INT32:
  110. _value = output.int_val
  111. elif output.dtype == tf_predict_pb2.DT_INT64:
  112. _value = output.int64_val
  113. elif output.dtype == tf_predict_pb2.DT_DOUBLE:
  114. _value = output.double_val
  115. elif output.dtype == tf_predict_pb2.DT_STRING:
  116. _value = output.string_val
  117. elif output.dtype == tf_predict_pb2.DT_BOOL:
  118. _value = output.bool_val
  119. return np.array(_value).reshape(response.outputs[output_name].array_shape.dim)
  120. def vpc_requests(url, authorization, request_data, list_outputs):
  121. headers = {"Authorization": authorization}
  122. dict_outputs = dict()
  123. response = tf_predict_pb2.PredictResponse()
  124. resp = requests.post(url, data=request_data, headers=headers)
  125. if resp.status_code != 200:
  126. print(resp.status_code, resp.content)
  127. # 从 common.logging 导入 log,避免依赖 common/Utils.py
  128. from BiddingKG.dl.common.logging import log
  129. log("调用pai-eas接口出错,authorization:" + str(authorization))
  130. return None
  131. else:
  132. response = tf_predict_pb2.PredictResponse()
  133. response.ParseFromString(resp.content)
  134. for _output in list_outputs:
  135. dict_outputs[_output] = get_values(response, _output)
  136. return dict_outputs