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

Function tensor__matmul__method

paddle/fluid/pybind/eager_math_op_patch.cc:1296–1437  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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
1308static 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)) {

Callers

nothing calls this directly

Calls 15

SetPythonStackFunction · 0.85
InstanceFunction · 0.85
CastPyArg2DoubleFunction · 0.85
ScalarClass · 0.85
InputsContainDistTensorFunction · 0.85
PyCheckTensorFunction · 0.85
CastPyArg2ScalarFunction · 0.85
GetExpectedPlaceMethod · 0.80
SetDeviceFunction · 0.70
PyCheckIntegerFunction · 0.70
IsNumpyTypeFunction · 0.70

Tested by

no test coverage detected