| 1294 | ConvertAllInputsToDistTensor(mesh, self_tensor, other_tensor); |
| 1295 | } |
| 1296 | } |
| 1297 | |
| 1298 | // 3. calculation |
| 1299 | VLOG(6) << "Calling remainder_ad_func in tensor__mod__method"; |
| 1300 | { |
| 1301 | eager_gil_scoped_release guard; |
| 1302 | ret = remainder_ad_func(self_tensor, other_tensor); |
| 1303 | } |
| 1304 | return ToPyObject(ret); |
| 1305 | EAGER_CATCH_AND_THROW_RETURN_NULL |
| 1306 | } |
| 1307 | |
| 1308 | static PyObject* tensor__rmod__method(TensorObject* self, |
| 1309 | PyObject* args, |
| 1310 | PyObject* kwargs) { |
| 1311 | phi::RecordEvent pythonc_record_event( |
| 1312 | "__rmod__ pybind_patch_func", phi::TracerEventType::UserDefined, 1); |
| 1313 | EAGER_TRY |
| 1314 | |
| 1315 | VLOG(6) << "Running Eager tensor__rmod__method"; |
| 1316 | |
| 1317 | SetPythonStack(); |
| 1318 | |
| 1319 | // Set Device ID |
| 1320 | auto place = egr::Controller::Instance().GetExpectedPlace(); |
| 1321 | SetDevice(place); |
| 1322 | |
| 1323 | Tensor ret; |
| 1324 | |
| 1325 | Tensor self_tensor = self->tensor; |
| 1326 | PyObject* other_obj = PyTuple_GET_ITEM(args, 0); |
| 1327 | |
| 1328 | // 1. scalar exists cases |
| 1329 | // there is no scalar_mod function for __rmod__ now |
| 1330 | if (PyFloat_Check(other_obj) || PyCheckInteger(other_obj) || |
| 1331 | IsNumpyType(other_obj)) { |
| 1332 | if (PyFloat_Check(other_obj)) { |
| 1333 | if (_supported_int_dtype_.find(self_tensor.dtype()) != |
| 1334 | _supported_int_dtype_.end()) { |
| 1335 | eager_gil_scoped_release guard; |
| 1336 | self_tensor = cast_ad_func(self_tensor, DataType::FLOAT32); |
| 1337 | } |
| 1338 | } else if (PyCheckInteger(other_obj) && |
| 1339 | self_tensor.dtype() == DataType::BOOL) { |
| 1340 | eager_gil_scoped_release guard; |
| 1341 | self_tensor = cast_ad_func(self_tensor, DataType::INT64); |
| 1342 | } |
| 1343 | } else if (PyComplex_Check(other_obj)) { |
| 1344 | if (is_support_complex(self_tensor.dtype()) == false) { |
| 1345 | eager_gil_scoped_release guard; |
| 1346 | self_tensor = cast_ad_func( |
| 1347 | self_tensor, promoteTypes(self_tensor.dtype(), DataType::COMPLEX64)); |
| 1348 | } |
| 1349 | } |
| 1350 | |
| 1351 | // 2. create or get tensor for other_obj |
| 1352 | Tensor other_tensor; |
| 1353 | if (PyCheckTensor(other_obj)) { |
nothing calls this directly
no test coverage detected