| 2603 | } |
| 2604 | |
| 2605 | std::vector<phi::Scalar> CastPyArg2ScalarArray(PyObject* obj, |
| 2606 | const std::string& op_type, |
| 2607 | ssize_t arg_pos) { |
| 2608 | if (obj == Py_None) { |
| 2609 | PADDLE_THROW(common::errors::InvalidType( |
| 2610 | "%s(): argument (position %d) must be " |
| 2611 | "a list of int, float, or bool, but got %s", |
| 2612 | op_type, |
| 2613 | arg_pos + 1, |
| 2614 | ((PyTypeObject*)obj->ob_type)->tp_name)); // NOLINT |
| 2615 | } |
| 2616 | |
| 2617 | PyTypeObject* type = obj->ob_type; |
| 2618 | auto type_name = std::string(type->tp_name); |
| 2619 | VLOG(4) << "type_name: " << type_name; |
| 2620 | if (PyList_Check(obj)) { |
| 2621 | Py_ssize_t len = PyList_Size(obj); |
| 2622 | if (len == 0) { |
| 2623 | return std::vector<phi::Scalar>({}); |
| 2624 | } |
| 2625 | PyObject* item = nullptr; |
| 2626 | item = PyList_GetItem(obj, 0); |
| 2627 | if (PyObject_CheckFloat(item)) { |
| 2628 | std::vector<phi::Scalar> value; |
| 2629 | for (Py_ssize_t i = 0; i < len; i++) { |
| 2630 | item = PyList_GetItem(obj, i); |
| 2631 | value.emplace_back(phi::Scalar{PyObject_ToDouble(item)}); |
| 2632 | } |
| 2633 | return value; |
| 2634 | } else if (PyObject_CheckLong(item)) { |
| 2635 | std::vector<phi::Scalar> value; |
| 2636 | for (Py_ssize_t i = 0; i < len; i++) { |
| 2637 | item = PyList_GetItem(obj, i); |
| 2638 | value.emplace_back(phi::Scalar{PyObject_ToInt64(item)}); |
| 2639 | } |
| 2640 | return value; |
| 2641 | } else if (PyObject_CheckComplexOrToComplex(&item)) { |
| 2642 | std::vector<phi::Scalar> value; |
| 2643 | for (Py_ssize_t i = 0; i < len; i++) { |
| 2644 | item = PyList_GetItem(obj, i); |
| 2645 | Py_complex v = PyComplex_AsCComplex(item); |
| 2646 | value.emplace_back(phi::Scalar{std::complex<double>(v.real, v.imag)}); |
| 2647 | } |
| 2648 | return value; |
| 2649 | } else { |
| 2650 | PADDLE_THROW(common::errors::InvalidType( |
| 2651 | "%s(): argument (position %d) must be " |
| 2652 | "a list of int, float, complex, or bool, but got %s", |
| 2653 | op_type, |
| 2654 | arg_pos + 1, |
| 2655 | ((PyTypeObject*)item->ob_type)->tp_name)); // NOLINT |
| 2656 | } |
| 2657 | } else { |
| 2658 | PADDLE_THROW(common::errors::InvalidType( |
| 2659 | "%s(): argument (position %d) must be " |
| 2660 | "a list of int, float, complex, or bool, but got %s", |
| 2661 | op_type, |
| 2662 | arg_pos + 1, |
nothing calls this directly
no test coverage detected