!30155 AllReduce input and output size aligned by 512 in ascend device context

Merge pull request !30155 from laiyongqiang/master_allreduce
This commit is contained in:
i-robot 2022-02-19 09:56:54 +00:00 committed by Gitee
commit 996e7bffaa
No known key found for this signature in database
GPG Key ID: 173E9B9CA92EEF8F
2 changed files with 17 additions and 10 deletions

View File

@ -26,6 +26,7 @@
#include "utils/utils.h"
#include "plugin/device/ascend/hal/device/kernel_select_ascend.h"
#include "runtime/device/kernel_adjust.h"
#include "runtime/device/memory_manager.h"
#include "plugin/device/ascend/hal/device/ascend_stream_assign.h"
#include "plugin/device/ascend/hal/device/kernel_build_ascend.h"
#include "plugin/device/ascend/hal/hardware/ascend_graph_optimization.h"
@ -525,8 +526,19 @@ void AscendDeviceContext::FreeMemory(DeviceAddress *const &address) const {
bool AscendDeviceContext::AllocateContinuousMemory(const std::vector<DeviceAddressPtr> &addr_list, size_t total_size,
const std::vector<size_t> &size_list) const {
MS_EXCEPTION_IF_NULL(runtime_instance_);
if (addr_list.size() != size_list.size()) {
MS_LOG(EXCEPTION) << "AllocateContinuousMemory check failed: input address list size " << addr_list.size()
<< " vs input size list size " << size_list.size();
}
runtime_instance_->SetContext();
return mem_manager_->MallocContinuousMemFromMemPool(addr_list, total_size, size_list);
size_t align_total_size = 0;
std::vector<size_t> align_size_list;
for (size_t i = 0; i < size_list.size(); i++) {
auto align_size = device::MemoryManager::GetCommonAlignSize(size_list[i]);
align_size_list.emplace_back(align_size);
align_total_size += align_size;
}
return mem_manager_->MallocContinuousMemFromMemPool(addr_list, align_total_size, align_size_list);
}
void *AscendDeviceContext::AllocateMemory(size_t size) const {

View File

@ -25,7 +25,6 @@
#include "mindrt/include/async/async.h"
#include "utils/log_adapter.h"
#include "utils/convert_utils.h"
#include "runtime/device/memory_manager.h"
namespace mindspore {
namespace runtime {
@ -76,10 +75,8 @@ void FetchContinuousMemoryInfo(const CNodePtr &node, std::vector<DeviceTensorPtr
for (size_t i = 0; i < intput_sizes.size(); ++i) {
const auto &device_tensor = AnfAlgo::GetPrevNodeMutableOutputAddr(node, i, false);
MS_EXCEPTION_IF_NULL(device_tensor);
auto origin_size = intput_sizes[i];
auto align_size = device::MemoryManager::GetCommonAlignSize(origin_size);
*total_size += align_size;
(void)size_list->emplace_back(align_size);
*total_size += intput_sizes[i];
(void)size_list->emplace_back(intput_sizes[i]);
(void)addr_list->emplace_back(device_tensor);
}
} else {
@ -87,10 +84,8 @@ void FetchContinuousMemoryInfo(const CNodePtr &node, std::vector<DeviceTensorPtr
for (size_t i = 0; i < output_sizes.size(); ++i) {
const auto &device_tensor = AnfAlgo::GetMutableOutputAddr(node, i, false);
MS_EXCEPTION_IF_NULL(device_tensor);
auto origin_size = output_sizes[i];
auto align_size = device::MemoryManager::GetCommonAlignSize(origin_size);
*total_size += align_size;
(void)size_list->emplace_back(align_size);
*total_size += output_sizes[i];
(void)size_list->emplace_back(output_sizes[i]);
(void)addr_list->emplace_back(device_tensor);
}
}