| 57 | } |
| 58 | |
| 59 | DataType InferScalarType(PyObject* object) { |
| 60 | if (PyBool_Check(object)) { |
| 61 | return DataType::kBool; |
| 62 | } else if (PyLong_Check(object)) { |
| 63 | return DataType::kInt64; |
| 64 | } else if (PyArray_Check(object)) { |
| 65 | return numpy::GetOFDataTypeFromNpArray(reinterpret_cast<PyArrayObject*>(object)).GetOrThrow(); |
| 66 | } else if (PyArray_CheckScalar(object)) { |
| 67 | return numpy::NumpyTypeToOFDataType(PyArray_DescrFromScalar(object)->type_num).GetOrThrow(); |
| 68 | } else if (PySequence_Check(object)) { |
| 69 | int64_t length = PySequence_Length(object); |
| 70 | if (length == 0) { return DataType::kInt64; } |
| 71 | DataType scalar_type = DataType::kInvalidDataType; |
| 72 | for (int64_t i = 0; i < length; ++i) { |
| 73 | PyObjectPtr item(PySequence_GetItem(object, i)); |
| 74 | const auto& item_scalar_type = InferScalarType(item.get()); |
| 75 | if (scalar_type != DataType::kInvalidDataType) { |
| 76 | CHECK_EQ_OR_THROW(scalar_type, item_scalar_type) |
| 77 | << "Different scalar types are not allowed."; |
| 78 | } else { |
| 79 | scalar_type = item_scalar_type; |
| 80 | } |
| 81 | } |
| 82 | return scalar_type; |
| 83 | } |
| 84 | THROW(TypeError) << "Can't infer scalar type of " << Py_TYPE(object)->tp_name; |
| 85 | return DataType::kInvalidDataType; |
| 86 | } |
| 87 | |
| 88 | void ParseScalar(PyObject* object, char* data, const DataType& dtype) { |
| 89 | if (dtype == DataType::kInt64) { |
no test coverage detected