MCPcopy Create free account
hub / github.com/PaddlePaddle/Paddle / CastPyArg2Scalars

Function CastPyArg2Scalars

paddle/fluid/pybind/op_function_common.cc:1045–1101  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

1043}
1044
1045std::vector<paddle::experimental::Scalar> CastPyArg2Scalars(
1046 PyObject* obj, const std::string& op_type, ssize_t arg_pos) {
1047 if (obj == Py_None) {
1048 PADDLE_THROW(common::errors::InvalidType(
1049 "%s(): argument (position %d) must be "
1050 "a list of int, float, or bool, but got %s",
1051 op_type,
1052 arg_pos + 1,
1053 ((PyTypeObject*)obj->ob_type)->tp_name)); // NOLINT
1054 }
1055
1056 PyTypeObject* type = obj->ob_type;
1057 auto type_name = std::string(type->tp_name);
1058 VLOG(4) << "type_name: " << type_name;
1059 if (PyList_Check(obj)) {
1060 Py_ssize_t len = PyList_Size(obj);
1061 PyObject* item = nullptr;
1062 item = PyList_GetItem(obj, 0);
1063 if (PyObject_CheckFloat(item)) {
1064 std::vector<paddle::experimental::Scalar> value;
1065 for (Py_ssize_t i = 0; i < len; i++) {
1066 item = PyList_GetItem(obj, i);
1067 value.emplace_back(
1068 paddle::experimental::Scalar{PyObject_ToDouble(item)});
1069 }
1070 return value;
1071 } else if (PyObject_CheckLong(item)) {
1072 std::vector<paddle::experimental::Scalar> value;
1073 for (Py_ssize_t i = 0; i < len; i++) {
1074 item = PyList_GetItem(obj, i);
1075 value.emplace_back(
1076 paddle::experimental::Scalar{PyObject_ToInt64(item)});
1077 }
1078 return value;
1079 } else if (PyObject_CheckComplexOrToComplex(&item)) {
1080 std::vector<paddle::experimental::Scalar> value;
1081 for (Py_ssize_t i = 0; i < len; i++) {
1082 item = PyList_GetItem(obj, i);
1083 Py_complex v = PyComplex_AsCComplex(item);
1084 value.emplace_back(
1085 paddle::experimental::Scalar{std::complex<double>(v.real, v.imag)});
1086 }
1087 return value;
1088 }
1089 } else {
1090 PADDLE_THROW(common::errors::InvalidType(
1091 "%s(): argument (position %d) must be "
1092 "a list of int, float, complex, or bool, but got %s",
1093 op_type,
1094 arg_pos + 1,
1095 ((PyTypeObject*)obj->ob_type)->tp_name)); // NOLINT
1096 }
1097
1098 // Fake a ScalarArray
1099 return std::vector<paddle::experimental::Scalar>(
1100 {paddle::experimental::Scalar(1.0)});
1101}
1102

Callers 1

CastPyArg2AttrScalarsFunction · 0.85

Calls 7

PyObject_CheckFloatFunction · 0.85
PyObject_ToDoubleFunction · 0.85
PyObject_CheckLongFunction · 0.85
PyObject_ToInt64Function · 0.85
ScalarClass · 0.85
emplace_backMethod · 0.45

Tested by

no test coverage detected