From 7c8ff29bf934076389e45287ef15fed8d302e515 Mon Sep 17 00:00:00 2001 From: chendongsheng Date: Fri, 16 Jul 2021 23:08:46 +0800 Subject: [PATCH] fixed codex warning --- mindspore/ccsrc/ps/constants.h | 16 +++ mindspore/ccsrc/ps/core/abstract_node.cc | 4 +- mindspore/ccsrc/ps/core/abstract_node.h | 10 +- .../ccsrc/ps/core/communicator/http_client.cc | 4 +- .../ccsrc/ps/core/communicator/http_client.h | 1 + .../core/communicator/http_message_handler.cc | 2 - .../ccsrc/ps/core/communicator/tcp_client.h | 4 +- .../ccsrc/ps/core/communicator/tcp_server.h | 4 +- mindspore/ccsrc/ps/core/node.cc | 6 +- mindspore/ccsrc/ps/core/node_manager.cc | 6 +- mindspore/ccsrc/ps/core/scheduler_node.cc | 6 +- mindspore/ccsrc/ps/parameter_server.cc | 13 ++- mindspore/ccsrc/ps/ps_context.cc | 8 +- mindspore/ccsrc/ps/worker.cc | 103 +++++++++--------- 14 files changed, 102 insertions(+), 85 deletions(-) diff --git a/mindspore/ccsrc/ps/constants.h b/mindspore/ccsrc/ps/constants.h index b6f953d5747..4d0ed5bb960 100644 --- a/mindspore/ccsrc/ps/constants.h +++ b/mindspore/ccsrc/ps/constants.h @@ -75,6 +75,22 @@ constexpr int64_t kPullCmd = 51; constexpr size_t kInvalidKey = UINT64_MAX; constexpr int64_t kInvalidID = -1; +constexpr int64_t kGradIndex = 0; +constexpr int64_t kIndiceIndex = 1; +constexpr int64_t kFirstDimSize = 2; +constexpr int64_t kOutDimSize = 3; + +constexpr int64_t kBase = 10; +constexpr float kStdDev = 0.01; + +constexpr int64_t kSparseLazyAdamIndex = 2; +constexpr int64_t kSparseFtrlIndex = 3; +constexpr int64_t kSparseGradIndex = 6; +constexpr int64_t kSparseIndiceIndex = 7; + +constexpr int64_t kHeartbeatTimes = 2; +constexpr int64_t kGradValue = -100; + constexpr uint32_t kMaxMessageSize = static_cast(100 * (uint32_t(1) << 20)); constexpr char kServerNum[] = "server_num"; constexpr char kWorkerNum[] = "worker_num"; diff --git a/mindspore/ccsrc/ps/core/abstract_node.cc b/mindspore/ccsrc/ps/core/abstract_node.cc index c5b5dc5bc76..e5e3f27fd52 100644 --- a/mindspore/ccsrc/ps/core/abstract_node.cc +++ b/mindspore/ccsrc/ps/core/abstract_node.cc @@ -338,8 +338,8 @@ std::pair AbstractNode::CollectiveReceiveAsync(const NodeRol uint64_t rank_request_id = NextExpectedRankRequestId(rank_id); receive_messages_done_[std::make_pair(rank_id, rank_request_id)] = false; if (received_data_.count(std::make_pair(rank_id, rank_request_id)) > 0) { - auto res = received_data_[std::make_pair(rank_id, rank_request_id)]; - *output = res; + auto res_output = received_data_[std::make_pair(rank_id, rank_request_id)]; + *output = res_output; received_data_.erase(std::make_pair(rank_id, rank_request_id)); receive_messages_done_[std::make_pair(rank_id, rank_request_id)] = true; MS_LOG(DEBUG) << "Receive data from rank id:" << rank_id << ", the rank request id is:" << rank_request_id; diff --git a/mindspore/ccsrc/ps/core/abstract_node.h b/mindspore/ccsrc/ps/core/abstract_node.h index 8918f919651..fafca84c0e3 100644 --- a/mindspore/ccsrc/ps/core/abstract_node.h +++ b/mindspore/ccsrc/ps/core/abstract_node.h @@ -62,7 +62,7 @@ class AbstractNode : public Node { using DataPtr = std::shared_ptr; using VectorPtr = std::shared_ptr>; - bool Broadcast(const enum NodeRole &node_role, const DataPtr &message, size_t size, int command, + bool Broadcast(const NodeRole &node_role, const DataPtr &message, size_t size, int command, const uint32_t &timeout = kCommTimeoutInSeconds); // When the business layer finish scale out, it should call this function @@ -84,18 +84,18 @@ class AbstractNode : public Node { // Set the callback corresponding to the custom event. void RegisterCustomEventCallback(const uint32_t &event, const EventCallback &event_cb); - bool Send(const enum NodeRole &node_role, const uint32_t &rank_id, const DataPtr &data, size_t len, int command, + bool Send(const NodeRole &node_role, const uint32_t &rank_id, const DataPtr &data, size_t len, int command, const uint32_t &timeout = kCommTimeoutInSeconds); bool Send(const NodeRole &node_role, const std::vector &rank_ids, const std::vector &data, const std::vector &lens, int command, const uint32_t &timeout = kCommTimeoutInSeconds); - bool Send(const enum NodeRole &node_role, const uint32_t &rank_id, const DataPtr &message, size_t len, int command, + bool Send(const NodeRole &node_role, const uint32_t &rank_id, const DataPtr &message, size_t len, int command, VectorPtr *output, const uint32_t &timeout = kCommTimeoutInSeconds); bool Send(const NodeRole &node_role, const std::vector &rank_ids, const std::vector &data, const std::vector &data_lens, int command, std::vector *output, const uint32_t &timeout = kCommTimeoutInSeconds); - uint64_t CollectiveSendAsync(const enum NodeRole &node_role, const uint32_t &rank_id, const void *data, size_t size); - std::pair CollectiveReceiveAsync(const enum NodeRole &node_role, const uint32_t &rank_id, + uint64_t CollectiveSendAsync(const NodeRole &node_role, const uint32_t &rank_id, const void *data, size_t size); + std::pair CollectiveReceiveAsync(const NodeRole &node_role, const uint32_t &rank_id, VectorPtr *output); bool CollectiveWait(const std::pair &request_id, const uint32_t &timeout = kCommTimeoutInSeconds); diff --git a/mindspore/ccsrc/ps/core/communicator/http_client.cc b/mindspore/ccsrc/ps/core/communicator/http_client.cc index 056e5542af3..18706a221b6 100644 --- a/mindspore/ccsrc/ps/core/communicator/http_client.cc +++ b/mindspore/ccsrc/ps/core/communicator/http_client.cc @@ -112,7 +112,7 @@ int HttpClient::ReadHeaderDoneCallback(struct evhttp_request *request, void *arg MS_LOG(DEBUG) << "The key:" << header->key << ",The value:" << header->value; std::string len = "Content-Length"; if (!strcmp(header->key, len.c_str())) { - handler->set_content_len(strtouq(header->value, nullptr, 10)); + handler->set_content_len(strtouq(header->value, nullptr, kBase)); handler->InitBodySize(); } } @@ -132,7 +132,7 @@ void HttpClient::ReadChunkDataCallback(struct evhttp_request *request, void *arg } } -void HttpClient::RequestErrorCallback(enum evhttp_request_error error, void *arg) { +void HttpClient::RequestErrorCallback(evhttp_request_error error, void *arg) { MS_EXCEPTION_IF_NULL(arg); auto handler = static_cast(arg); MS_LOG(ERROR) << "The request failed, the error is:" << error; diff --git a/mindspore/ccsrc/ps/core/communicator/http_client.h b/mindspore/ccsrc/ps/core/communicator/http_client.h index 859b63dfb2c..bc47ac1e454 100644 --- a/mindspore/ccsrc/ps/core/communicator/http_client.h +++ b/mindspore/ccsrc/ps/core/communicator/http_client.h @@ -42,6 +42,7 @@ #include "ps/core/communicator/http_message_handler.h" #include "ps/core/comm_util.h" #include "utils/convert_utils_base.h" +#include "ps/constants.h" namespace mindspore { namespace ps { diff --git a/mindspore/ccsrc/ps/core/communicator/http_message_handler.cc b/mindspore/ccsrc/ps/core/communicator/http_message_handler.cc index 0363f8217fd..1be9be96445 100644 --- a/mindspore/ccsrc/ps/core/communicator/http_message_handler.cc +++ b/mindspore/ccsrc/ps/core/communicator/http_message_handler.cc @@ -29,8 +29,6 @@ #include #include #include -#include -#include #include #include diff --git a/mindspore/ccsrc/ps/core/communicator/tcp_client.h b/mindspore/ccsrc/ps/core/communicator/tcp_client.h index c73b1e7666b..c80e56352be 100644 --- a/mindspore/ccsrc/ps/core/communicator/tcp_client.h +++ b/mindspore/ccsrc/ps/core/communicator/tcp_client.h @@ -53,7 +53,7 @@ class TcpClient { std::function &, const Protos &, const void *, size_t size)>; using OnTimer = std::function; - explicit TcpClient(const std::string &address, std::uint16_t port, Configuration *config); + explicit TcpClient(const std::string &address, std::uint16_t port, Configuration *const config); virtual ~TcpClient(); std::string GetServerAddress() const; @@ -107,7 +107,7 @@ class TcpClient { std::atomic is_stop_; std::atomic is_connected_; // The Configuration file - Configuration *config_; + Configuration *const config_; }; } // namespace core } // namespace ps diff --git a/mindspore/ccsrc/ps/core/communicator/tcp_server.h b/mindspore/ccsrc/ps/core/communicator/tcp_server.h index 05bfd71640c..c43d4fdcbe8 100644 --- a/mindspore/ccsrc/ps/core/communicator/tcp_server.h +++ b/mindspore/ccsrc/ps/core/communicator/tcp_server.h @@ -86,7 +86,7 @@ class TcpServer { using OnTimerOnce = std::function; using OnTimer = std::function; - TcpServer(const std::string &address, std::uint16_t port, Configuration *config); + TcpServer(const std::string &address, std::uint16_t port, Configuration *const config); TcpServer(const TcpServer &server); virtual ~TcpServer(); @@ -142,7 +142,7 @@ class TcpServer { OnTimerOnce on_timer_once_callback_; OnTimer on_timer_callback_; // The Configuration file - Configuration *config_; + Configuration *const config_; }; } // namespace core } // namespace ps diff --git a/mindspore/ccsrc/ps/core/node.cc b/mindspore/ccsrc/ps/core/node.cc index 42fd76444ff..c7f92e4106f 100644 --- a/mindspore/ccsrc/ps/core/node.cc +++ b/mindspore/ccsrc/ps/core/node.cc @@ -32,11 +32,11 @@ std::string Node::BoundIp() const { return node_info_.ip_; } bool Node::WaitForStart(const uint32_t &timeout) { std::unique_lock lock(wait_start_mutex_); bool res = wait_start_cond_.wait_for(lock, std::chrono::seconds(timeout), [this] { - bool res = this->is_ready_.load(); - if (res) { + bool result = this->is_ready_.load(); + if (result) { MS_LOG(INFO) << "The node id:" << node_info_.node_id_ << " is success start!"; } - return res; + return result; }); return res; } diff --git a/mindspore/ccsrc/ps/core/node_manager.cc b/mindspore/ccsrc/ps/core/node_manager.cc index 1760e055d37..453120ba86b 100644 --- a/mindspore/ccsrc/ps/core/node_manager.cc +++ b/mindspore/ccsrc/ps/core/node_manager.cc @@ -171,9 +171,9 @@ void NodeManager::UpdateCluster() { if (!timeout_nodes_info_.empty()) { UpdateClusterState(ClusterState::NODE_TIMEOUT); - for (auto it = timeout_nodes_info_.begin(); it != timeout_nodes_info_.end(); ++it) { - heartbeats_.erase(it->first); - finish_nodes_id_.insert(it->first); + for (auto iter = timeout_nodes_info_.begin(); iter != timeout_nodes_info_.end(); ++iter) { + heartbeats_.erase(iter->first); + finish_nodes_id_.insert(iter->first); } } diff --git a/mindspore/ccsrc/ps/core/scheduler_node.cc b/mindspore/ccsrc/ps/core/scheduler_node.cc index b79e372afea..432ab2c550b 100644 --- a/mindspore/ccsrc/ps/core/scheduler_node.cc +++ b/mindspore/ccsrc/ps/core/scheduler_node.cc @@ -405,7 +405,7 @@ void SchedulerNode::StartUpdateClusterStateTimer() { if (node_manager_.GetClusterState() == ClusterState::CLUSTER_EXIT) { std::this_thread::sleep_for( - std::chrono::seconds(PSContext::instance()->cluster_config().heartbeat_interval * 2)); + std::chrono::seconds(PSContext::instance()->cluster_config().heartbeat_interval * kHeartbeatTimes)); is_finish_ = true; wait_finish_cond_.notify_all(); } @@ -743,7 +743,9 @@ void SchedulerNode::StartRestfulServer(const std::string &address, std::uint16_t void SchedulerNode::StopRestfulServer() { MS_LOG(INFO) << "Scheduler stop https server."; - http_server_->Stop(); + if (!http_server_->Stop()) { + MS_LOG(WARNING) << "Stop http server failed."; + } if (restful_thread_ != nullptr && restful_thread_->joinable()) { restful_thread_->join(); } diff --git a/mindspore/ccsrc/ps/parameter_server.cc b/mindspore/ccsrc/ps/parameter_server.cc index 7fde5188583..1e63b09a9b6 100644 --- a/mindspore/ccsrc/ps/parameter_server.cc +++ b/mindspore/ccsrc/ps/parameter_server.cc @@ -42,8 +42,8 @@ void ParameterServer::Run(const FuncGraphPtr &func_graph) { } bool ParameterServer::Init(const FuncGraphPtr &func_graph) { - pserver_num_ = std::strtol(mindspore::common::GetEnv(kEnvPServerNum).c_str(), nullptr, 10); - worker_num_ = std::strtol(mindspore::common::GetEnv(kEnvWorkerNum).c_str(), nullptr, 10); + pserver_num_ = std::strtol(mindspore::common::GetEnv(kEnvPServerNum).c_str(), nullptr, kBase); + worker_num_ = std::strtol(mindspore::common::GetEnv(kEnvWorkerNum).c_str(), nullptr, kBase); func_graph_ = func_graph; handler_.reset(new ServerHandler(this)); handler_->Init(); @@ -176,10 +176,11 @@ void ParameterServer::InitEmbeddingTable( MS_EXCEPTION_IF_NULL(embedding); float *embedding_data = embedding->data(); std::default_random_engine engine; - std::normal_distribution random(0, 0.01); + std::normal_distribution random(0, kStdDev); if (ps::PsDataPrefetch::GetInstance().cache_enable()) { if (param_init_info.param_type_ == kWeight) { - InitRandomNormal(0, 0.01, input_shapes, param_init_info.global_seed_, param_init_info.op_seed_, embedding_data); + InitRandomNormal(0, kStdDev, input_shapes, param_init_info.global_seed_, param_init_info.op_seed_, + embedding_data); } else if (param_init_info.param_type_ == kAccumulation) { for (size_t i = 0; i < total_dims; i++) { embedding_data[i] = param_init_info.init_val_; @@ -259,7 +260,7 @@ void ParameterServer::UpdateWeights() { void ParameterServer::AccumGrad(const Keys &keys, const Values &values, const Lengths &lengths) { std::unique_lock lock(mutex_); const Key &key = keys[0]; - bool no_sparse_grad = values.size() == 1 && values[0] == -100; + bool no_sparse_grad = values.size() == 1 && values[0] == kGradValue; if (!no_sparse_grad) { std::shared_ptr optim_info = optim_infos_[key]; @@ -338,7 +339,7 @@ void ParameterServer::DoEmbeddingLookup(Key key, const LookupIds &lookup_ids, KV embedding_table->addr = table_ptr->data(); embedding_table->size = table_ptr->size() * sizeof(float); - std::unique_ptr tmp_ids(new int[lookup_ids.size()]); + std::unique_ptr tmp_ids = std::make_unique(lookup_ids.size()); MS_EXCEPTION_IF_NULL(tmp_ids); for (size_t i = 0; i < lookup_ids.size(); i++) { tmp_ids[i] = static_cast(lookup_ids[i]); diff --git a/mindspore/ccsrc/ps/ps_context.cc b/mindspore/ccsrc/ps/ps_context.cc index b1dce5f3bcd..110c17a3761 100644 --- a/mindspore/ccsrc/ps/ps_context.cc +++ b/mindspore/ccsrc/ps/ps_context.cc @@ -48,11 +48,11 @@ void PSContext::SetPSEnable(bool enabled) { MS_LOG(INFO) << "MS_ROLE is " << ms_role; } - worker_num_ = std::strtol(common::GetEnv(kEnvWorkerNum).c_str(), nullptr, 10); - server_num_ = std::strtol(common::GetEnv(kEnvPServerNum).c_str(), nullptr, 10); + worker_num_ = std::strtol(common::GetEnv(kEnvWorkerNum).c_str(), nullptr, kBase); + server_num_ = std::strtol(common::GetEnv(kEnvPServerNum).c_str(), nullptr, kBase); scheduler_host_ = common::GetEnv(kEnvSchedulerHost); - scheduler_port_ = std::strtol(common::GetEnv(kEnvSchedulerPort).c_str(), nullptr, 10); - scheduler_manage_port_ = std::strtol(common::GetEnv(kEnvSchedulerManagePort).c_str(), nullptr, 10); + scheduler_port_ = std::strtol(common::GetEnv(kEnvSchedulerPort).c_str(), nullptr, kBase); + scheduler_manage_port_ = std::strtol(common::GetEnv(kEnvSchedulerManagePort).c_str(), nullptr, kBase); cluster_config_ = std::make_unique(worker_num_, server_num_, scheduler_host_, scheduler_port_); } else { MS_LOG(INFO) << "PS mode is disabled."; diff --git a/mindspore/ccsrc/ps/worker.cc b/mindspore/ccsrc/ps/worker.cc index f3ea3d509ed..2550a5aa2ce 100644 --- a/mindspore/ccsrc/ps/worker.cc +++ b/mindspore/ccsrc/ps/worker.cc @@ -62,19 +62,19 @@ void Worker::Push(const std::vector &keys, std::vector addrs, int64_t optim_id = key_to_optimId_[key]; MS_LOG(INFO) << "The key is:" << key << " the optim_id:" << optim_id; bool is_sparse = false; - if (optim_id == 1 || optim_id == 2 || optim_id == 3) { + if (optim_id == 1 || optim_id == kSparseLazyAdamIndex || optim_id == kSparseFtrlIndex) { is_sparse = true; } int64_t grad_index = -1; int64_t indice_index = -1; // Sparse adam gradient - if (optim_id == 1 || optim_id == 2) { - grad_index = 6; - indice_index = 7; + if (optim_id == 1 || optim_id == kSparseLazyAdamIndex) { + grad_index = kSparseGradIndex; + indice_index = kSparseIndiceIndex; // Sparse ftrl gradient - } else if (optim_id == 3) { + } else if (optim_id == kSparseFtrlIndex) { grad_index = 0; indice_index = 1; } @@ -87,9 +87,9 @@ void Worker::Push(const std::vector &keys, std::vector addrs, void *src_data = reinterpret_cast(addrs[i]); MS_EXCEPTION_IF_NULL(dst_data); MS_EXCEPTION_IF_NULL(src_data); - int size = sizes[i] * sizeof(float); - size_t dest_size = IntToSize(size); - size_t src_size = IntToSize(size); + size_t size = sizes[i] * sizeof(float); + size_t dest_size = size; + size_t src_size = size; auto ret = memcpy_s(dst_data, dest_size, src_data, src_size); if (ret != 0) { MS_LOG(EXCEPTION) << "memcpy_s error, errorno(" << ret << ")"; @@ -114,8 +114,8 @@ void Worker::Push(const std::vector &keys, std::vector addrs, MS_LOG(DEBUG) << "The keys:" << keys << " the total_buffer:" << total_buffer << " the sizes_int:" << sizes_int << " the grad_index:" << grad_index << " the indice_index:" << indice_index << " the first_dim_size:" << first_dim_size << " the outer_dim_size" << outer_dim_size; - PushSparseData(std::vector(keys), total_buffer, std::vector(sizes_int), grad_index, indice_index, - first_dim_size, outer_dim_size); + PushSparseData(std::vector(keys), total_buffer, std::vector(sizes_int), LongToSize(grad_index), + LongToSize(indice_index), LongToSize(first_dim_size), LongToSize(outer_dim_size)); } } @@ -191,7 +191,7 @@ void Worker::AddEmbeddingTable(const Key &key, const size_t &row_count) { uint64_t begin = 0; uint64_t end = 0; for (int64_t i = 0; i < server_num_; i++) { - int64_t local_row_cnt = Util::LocalShard(row_count, i, server_num_); + size_t local_row_cnt = LongToSize(Util::LocalShard(row_count, i, server_num_)); MS_LOG(DEBUG) << "The row_count:" << row_count << " the local_row_cnt:" << local_row_cnt; if (i == 0) { end = local_row_cnt - 1; @@ -322,16 +322,16 @@ void Worker::DoPSEmbeddingLookup(const Key &key, const std::vector &lookup_ values->push_back(message.values(j)); } for (auto k = 0; k < message.keys_size(); k++) { - const Key &key = message.keys(k); - keys->push_back(key); + const Key &message_key = message.keys(k); + keys->push_back(message_key); } } for (size_t i = 0; i < keys->size(); i++) { - const Key &key = keys->at(i); + const Key &map_key = keys->at(i); float *addr = values->data() + value_offset; value_offset += single_id_len; - id_addr_map[key] = std::make_shared>(std::make_pair(addr, single_id_len)); + id_addr_map[map_key] = std::make_shared>(std::make_pair(addr, single_id_len)); } float *result_addr = lookup_result->data(); @@ -346,18 +346,18 @@ void Worker::DoPSEmbeddingLookup(const Key &key, const std::vector &lookup_ offset += single_id_len; continue; } - const Key &key = static_cast(lookup_ids[i]); - auto &pair = id_addr_map[key]; - int64_t size = single_id_len * sizeof(float); + const Key &id_key = static_cast(lookup_ids[i]); + auto &pair = id_addr_map[id_key]; + size_t size = LongToSize(single_id_len * sizeof(float)); dst_size = size; src_size = size; dst_data = result_addr + offset; src_data = pair->first; MS_EXCEPTION_IF_NULL(dst_data); MS_EXCEPTION_IF_NULL(src_data); - auto ret = memcpy_s(dst_data, dst_size, src_data, src_size); - if (ret != 0) { - MS_LOG(EXCEPTION) << "memcpy_s error, errorno(" << ret << ")"; + auto mem_ret = memcpy_s(dst_data, dst_size, src_data, src_size); + if (mem_ret != 0) { + MS_LOG(EXCEPTION) << "memcpy_s error, errorno(" << mem_ret << ")"; return; } offset += single_id_len; @@ -532,13 +532,13 @@ bool Worker::IsReadyForPull(const Key &key) { } } -void Worker::PrepareSparseGradient(const size_t begin, const size_t end, const std::unordered_set &distinct_ids, +void Worker::PrepareSparseGradient(const size_t, const size_t, const std::unordered_set &distinct_ids, const std::vector> &indice_to_grads, const int *all_indice, const size_t segment_size, float *gradient, int *indices) { MS_EXCEPTION_IF_NULL(all_indice); MS_EXCEPTION_IF_NULL(gradient); MS_EXCEPTION_IF_NULL(indices); - int64_t offset = 0; + size_t offset = 0; int64_t index = 0; size_t segment_data_size = segment_size * sizeof(float); size_t dst_size; @@ -580,16 +580,16 @@ void Worker::BuildSparseValue(const std::vector &lengths, const size_t grad void *src_data = nullptr; for (size_t i = 0; i < lengths.size(); i++) { if (i != grad_index && i != indice_index) { - int data_size = lengths[i] * sizeof(float); + size_t data_size = lengths[i] * sizeof(float); dst_size = data_size; src_size = data_size; dst_data = reduced_data->data() + offset; src_data = const_cast(original_data) + offset; MS_EXCEPTION_IF_NULL(dst_data); MS_EXCEPTION_IF_NULL(src_data); - auto ret = memcpy_s(dst_data, dst_size, src_data, src_size); - if (ret != 0) { - MS_LOG(EXCEPTION) << "memcpy_s error, errorno(" << ret << ")"; + auto mem_ret = memcpy_s(dst_data, dst_size, src_data, src_size); + if (mem_ret != 0) { + MS_LOG(EXCEPTION) << "memcpy_s error, errorno(" << mem_ret << ")"; return; } } @@ -601,13 +601,12 @@ void Worker::BuildSparseValue(const std::vector &lengths, const size_t grad for (size_t i = 0; i < grad_index; i++) { grad_offset += lengths[i]; } - int64_t data_size = lengths[grad_index] * sizeof(float); + size_t data_size = lengths[grad_index] * sizeof(float); dst_size = data_size; src_size = data_size; dst_data = reduced_data->data() + grad_offset; src_data = const_cast(grads); MS_EXCEPTION_IF_NULL(dst_data); - MS_EXCEPTION_IF_NULL(src_data); auto ret = memcpy_s(dst_data, dst_size, src_data, src_size); if (ret != 0) { MS_LOG(EXCEPTION) << "memcpy_s error, errorno(" << ret << ")"; @@ -726,24 +725,25 @@ void Worker::SparsePartitioner(const KVMessage &send, PartitionKVMessages *parti // Init variables float *data = const_cast(send.values().data()); - if (attrs.count(0) == 0 || attrs.count(1) == 0 || attrs.count(2) == 0 || attrs.count(3) == 0) { + if (attrs.count(kGradIndex) == 0 || attrs.count(kIndiceIndex) == 0 || attrs.count(kFirstDimSize) == 0 || + attrs.count(kOutDimSize) == 0) { MS_LOG(EXCEPTION) << "Invalid attrs keys"; } - auto iter = attrs.find(0); + auto iter = attrs.find(kGradIndex); size_t grad_index = static_cast(iter->second); - iter = attrs.find(1); + iter = attrs.find(kIndiceIndex); size_t indice_index = static_cast(iter->second); - iter = attrs.find(2); + iter = attrs.find(kFirstDimSize); size_t first_dim_size = static_cast(iter->second); - iter = attrs.find(3); + iter = attrs.find(kOutDimSize); size_t outer_dim_size = static_cast(iter->second); - int grad_size = send.len()[grad_index]; - int indice_size = send.len()[indice_index]; - int segment_size = grad_size / indice_size; + size_t grad_size = send.len()[grad_index]; + size_t indice_size = send.len()[indice_index]; + size_t segment_size = grad_size / indice_size; - int64_t grad_offset = 0; - int64_t indice_offset = 0; + size_t grad_offset = 0; + size_t indice_offset = 0; for (size_t i = 0; i < grad_index; i++) { grad_offset += send.len()[i]; } @@ -757,7 +757,7 @@ void Worker::SparsePartitioner(const KVMessage &send, PartitionKVMessages *parti // Build the mappings of indice to gradient std::vector> indice_to_grads; - for (int i = 0; i < indice_size; i++) { + for (size_t i = 0; i < indice_size; i++) { int indice = indice_data[i]; float *grad = grad_data + i * segment_size; indice_to_grads.push_back(std::make_pair(indice, grad)); @@ -779,7 +779,7 @@ void Worker::SparsePartitioner(const KVMessage &send, PartitionKVMessages *parti // Prepare the sparse gradient and indice std::vector indice_ids; std::unordered_set distinct_ids; - for (int j = 0; j < indice_size; j++) { + for (size_t j = 0; j < indice_size; j++) { size_t indice = static_cast(indice_data[j]); if (indice >= begin && indice <= end) { indice_ids.push_back(indice); @@ -788,7 +788,7 @@ void Worker::SparsePartitioner(const KVMessage &send, PartitionKVMessages *parti } size_t indices_size = indice_ids.size(); if (indices_size > 0) { - int partition_segment_size = indices_size * segment_size; + size_t partition_segment_size = indices_size * segment_size; std::vector src_grad_data(partition_segment_size); std::vector src_indice_data(indices_size); PrepareSparseGradient(begin, end, distinct_ids, indice_to_grads, indice_data, segment_size, src_grad_data.data(), @@ -802,8 +802,7 @@ void Worker::SparsePartitioner(const KVMessage &send, PartitionKVMessages *parti first_dim_size, outer_dim_size, &unique_sparse_grad); // Update the length of reduce sparse gradient and indice - std::vector reduced_lens; - reduced_lens = {kvs.len().begin(), kvs.len().end()}; + std::vector reduced_lens = {kvs.len().begin(), kvs.len().end()}; reduced_lens[grad_index] = unique_sparse_grad.indices_size_ * segment_size; reduced_lens[indice_index] = unique_sparse_grad.indices_size_; @@ -822,7 +821,7 @@ void Worker::SparsePartitioner(const KVMessage &send, PartitionKVMessages *parti std::vector no_vals; std::vector no_lens; no_keys.push_back(key); - no_vals.push_back(-100); + no_vals.push_back(kGradValue); *kvs.mutable_values() = {no_vals.begin(), no_vals.end()}; *kvs.mutable_len() = {no_lens.begin(), no_lens.end()}; } @@ -833,14 +832,14 @@ void Worker::SparsePartitioner(const KVMessage &send, PartitionKVMessages *parti void Worker::RoundRobinPartitioner(const KVMessage &send, PartitionKVMessages *partition, const std::map &) { MS_EXCEPTION_IF_NULL(partition); - partition->resize(server_num_); + partition->resize(LongToSize(server_num_)); auto keys = send.keys(); auto values = send.values(); auto lens = send.len(); MS_LOG(INFO) << "the key size is:" << send.keys_size() << " the values size is:" << send.values_size() << " the lens:" << send.len_size(); - int64_t len; + size_t len; Key param_key; for (int i = 0; i < send.keys_size(); i++) { param_key = keys[i]; @@ -868,12 +867,12 @@ void Worker::RoundRobinPartitioner(const KVMessage &send, PartitionKVMessages *p void Worker::WorkerInitEmbeddingPartitioner(const KVMessage &send, std::vector> *partition, const std::map &attrs) { MS_EXCEPTION_IF_NULL(partition); - partition->resize(server_num_); + partition->resize(LongToSize(server_num_)); auto keys = send.keys(); auto values = send.values(); auto lens = send.len(); - size_t col_cnt = lens[0] / embedding_row_cnt_[keys[0]]; + int32_t col_cnt = lens[0] / embedding_row_cnt_[keys[0]]; const std::vector &ranges = *(embedding_table_ranges_[keys[0]]); for (size_t i = 0; i < ranges.size(); i++) { size_t offset_begin = ranges[i].begin() * col_cnt; @@ -887,7 +886,7 @@ void Worker::WorkerInitEmbeddingPartitioner(const KVMessage &send, std::vector &attrs) { + const std::map &) { MS_EXCEPTION_IF_NULL(partition); const float *embedding_vals = send.values().data(); const uint64_t *lookup_ids = send.len().data(); @@ -926,8 +925,8 @@ void Worker::UpdateEmbeddingPartitioner(const KVMessage &send, PartitionKVMessag void Worker::BroadcastPartitioner(const KVMessage &send, PartitionKVMessages *partition, const std::map &) { MS_EXCEPTION_IF_NULL(partition); - partition->resize(server_num_); - for (int64_t i = 0; i < server_num_; i++) { + partition->resize(LongToSize(server_num_)); + for (size_t i = 0; i < LongToSize(server_num_); i++) { partition->at(i).first = true; partition->at(i).second = send; }