| 1043 | } |
| 1044 | |
| 1045 | std::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 |
no test coverage detected