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

Function CastPyArg2ScalarArray

paddle/fluid/pybind/eager_utils.cc:2605–2665  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

2603}
2604
2605std::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,

Callers

nothing calls this directly

Calls 6

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

Tested by

no test coverage detected