| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154 |
- # -*- coding: utf-8 -*-
- """PAI-EAS 外部推理服务客户端。
- 按 ARCHITECTURE.md Phase 3 拆分建议,从 ``common/Utils.py`` 迁出。
- 类型:INFRA(外部服务适配)。
- 原位置:``common/Utils.py`` 中以下内容:
- - ``USE_PAI_EAS`` / ``API_URL`` / ``USE_API`` 全局变量(Phase 1 已改为从 infra/config 读取)
- - PAI-EAS 各模型的 authorization / url 常量
- - ``limitRun`` — TF session 分批推理
- - ``get_values`` — PAI-EAS 响应解析
- - ``vpc_requests`` — PAI-EAS VPC 请求
- ``common/Utils.py`` 仍 re-export 以上全部名称,老 import 不受影响。
- 注意:
- ``tf_predict_pb2`` 原代码中被注释掉(``# from pai_tf_predict_proto import tf_predict_pb2``),
- 本文件保持相同行为:尝试 import,失败则设为 None,``get_values`` / ``vpc_requests``
- 在调用时会因 ``tf_predict_pb2 is None`` 而报错,与原实现行为一致。
- """
- from __future__ import absolute_import
- import numpy as np
- import requests
- # PAI-EAS / 内部 API 配置走 infra/config(Phase 1)
- try:
- from BiddingKG.dl.infra import config as _infra_config
- _use_pai, _api_url, _use_api = _infra_config.pai_eas_settings()
- USE_PAI_EAS = _use_pai
- API_URL = _api_url or "http://127.0.0.1:888"
- USE_API = _use_api
- except Exception:
- # infra 不可用时回退默认值,保证其他功能不挂
- USE_PAI_EAS = False
- API_URL = "http://127.0.0.1:888"
- USE_API = False
- # pai_tf_predict_proto 可能在运行环境不可用,保持与原 Utils.py 一致的延迟/容错行为
- try:
- from pai_tf_predict_proto import tf_predict_pb2
- except ImportError:
- tf_predict_pb2 = None
- __all__ = [
- "USE_PAI_EAS",
- "API_URL",
- "USE_API",
- "tf_predict_pb2",
- "selffool_authorization",
- "selffool_url",
- "selffool_seg_authorization",
- "selffool_seg_url",
- "codename_authorization",
- "codename_url",
- "form_item_authorization",
- "form_item_url",
- "person_authorization",
- "person_url",
- "role_authorization",
- "role_url",
- "money_authorization",
- "money_url",
- "codeclasses_authorization",
- "codeclasses_url",
- "limitRun",
- "get_values",
- "vpc_requests",
- ]
- selffool_authorization = "NjlhMWFjMjVmNWYyNzI0MjY1OGQ1M2Y0ZmY4ZGY0Mzg3Yjc2MTVjYg=="
- selffool_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/selffool_gpu"
- selffool_seg_authorization = "OWUwM2Q0ZmE3YjYxNzU4YzFiMjliNGVkMTA3MzJkNjQ2MzJiYzBhZg=="
- selffool_seg_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/selffool_seg_gpu"
- codename_authorization = "Y2M5MDUxMzU1MTU4OGM3ZDk2ZmEzYjkxYmYyYzJiZmUyYTgwYTg5NA=="
- codename_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/codename_gpu"
- form_item_authorization = "ODdkZWY1YWY0NmNhNjU2OTI2NWY4YmUyM2ZlMDg1NTZjOWRkYTVjMw=="
- form_item_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/form"
- person_authorization = "N2I2MDU2N2Q2MGQ0ZWZlZGM3NDkyNTA1Nzc4YmM5OTlhY2MxZGU1Mw=="
- person_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/person"
- role_authorization = "OWM1ZDg5ZDEwYTEwYWI4OGNjYmRlMmQ1NzYwNWNlZGZkZmRmMjE4OQ=="
- role_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/role"
- money_authorization = "MDQyNjc2ZDczYjBhYmM4Yzc4ZGI4YjRmMjc3NGI5NTdlNzJiY2IwZA=="
- money_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/money"
- codeclasses_authorization = "MmUyNWIxZjQ2NjAzMWJlMGIzYzkxMjMzNWY5OWI3NzJlMWQ1ZjY4Yw=="
- codeclasses_url = "http://pai-eas-vpc.cn-beijing.aliyuncs.com/api/predict/codeclasses"
- def limitRun(sess, list_output, feed_dict, MAX_BATCH=1024):
- len_sample = 0
- if len(feed_dict.keys()) > 0:
- len_sample = len(feed_dict[list(feed_dict.keys())[0]])
- if len_sample > MAX_BATCH:
- list_result = [[] for _ in range(len(list_output))]
- _begin = 0
- while _begin < len_sample:
- new_dict = dict()
- for _key in feed_dict.keys():
- if isinstance(feed_dict[_key], (float, int, np.int32, np.float_, np.float16, np.float32, np.float64)):
- new_dict[_key] = feed_dict[_key]
- else:
- new_dict[_key] = feed_dict[_key][_begin:_begin + MAX_BATCH]
- _output = sess.run(list_output, feed_dict=new_dict)
- for _index in range(len(list_output)):
- list_result[_index].extend(_output[_index])
- _begin += MAX_BATCH
- else:
- list_result = sess.run(list_output, feed_dict=feed_dict)
- return list_result
- def get_values(response, output_name):
- """
- Get the value of a specified output tensor
- :param output_name: name of the output tensor
- :return: the content of the output tensor
- """
- output = response.outputs[output_name]
- if output.dtype == tf_predict_pb2.DT_FLOAT:
- _value = output.float_val
- elif output.dtype == tf_predict_pb2.DT_INT8 or output.dtype == tf_predict_pb2.DT_INT16 or \
- output.dtype == tf_predict_pb2.DT_INT32:
- _value = output.int_val
- elif output.dtype == tf_predict_pb2.DT_INT64:
- _value = output.int64_val
- elif output.dtype == tf_predict_pb2.DT_DOUBLE:
- _value = output.double_val
- elif output.dtype == tf_predict_pb2.DT_STRING:
- _value = output.string_val
- elif output.dtype == tf_predict_pb2.DT_BOOL:
- _value = output.bool_val
- return np.array(_value).reshape(response.outputs[output_name].array_shape.dim)
- def vpc_requests(url, authorization, request_data, list_outputs):
- headers = {"Authorization": authorization}
- dict_outputs = dict()
- response = tf_predict_pb2.PredictResponse()
- resp = requests.post(url, data=request_data, headers=headers)
- if resp.status_code != 200:
- print(resp.status_code, resp.content)
- # 从 common.logging 导入 log,避免依赖 common/Utils.py
- from BiddingKG.dl.common.logging import log
- log("调用pai-eas接口出错,authorization:" + str(authorization))
- return None
- else:
- response = tf_predict_pb2.PredictResponse()
- response.ParseFromString(resp.content)
- for _output in list_outputs:
- dict_outputs[_output] = get_values(response, _output)
- return dict_outputs
|