MCPcopy Create free account
hub / github.com/Oneflow-Inc/oneflow / InferScalarType

Function InferScalarType

oneflow/api/python/functional/indexing.cpp:59–86  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

57}
58
59DataType 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
88void ParseScalar(PyObject* object, char* data, const DataType& dtype) {
89 if (dtype == DataType::kInt64) {

Callers 1

ConvertToIndexingTensorFunction · 0.85

Calls 4

GetOFDataTypeFromNpArrayFunction · 0.85
NumpyTypeToOFDataTypeFunction · 0.85
GetOrThrowMethod · 0.45
getMethod · 0.45

Tested by

no test coverage detected