forked from huawei/mindspore2022
!25255 fix the weight update bug in the shared weight scene
Merge pull request !25255 from limingqi107/r1.5_lmq
This commit is contained in:
commit
b98b46cddc
|
|
@ -421,7 +421,6 @@ void DataPrepareActor::PrepareDataForWeightNode(const AnfNodePtr &backend_node,
|
|||
UpdateRefCount(host_tensor_address.get(), true);
|
||||
}
|
||||
MS_EXCEPTION_IF_NULL(host_tensor_address);
|
||||
DeviceTensorStore::GetInstance().Insert(front_node.get(), host_tensor_address);
|
||||
if (host_tensor_address->DeviceType() == device_tensor->DeviceType()) {
|
||||
AnfAlgo::SetOutputAddr(host_tensor_address, 0, backend_node.get());
|
||||
} else {
|
||||
|
|
@ -429,6 +428,9 @@ void DataPrepareActor::PrepareDataForWeightNode(const AnfNodePtr &backend_node,
|
|||
<< ", device tensor type:" << device_tensor->DeviceType();
|
||||
}
|
||||
}
|
||||
// Maybe the same host_tensor_address corresponds to the different front_node in shared weight scene,
|
||||
// so need update the device tensor store always.
|
||||
DeviceTensorStore::GetInstance().Insert(front_node.get(), host_tensor_address);
|
||||
|
||||
// If the ptr of device tensor is not nullptr, it indicates that the device data has been prepared.
|
||||
MS_EXCEPTION_IF_NULL(host_tensor_address);
|
||||
|
|
|
|||
Loading…
Reference in New Issue