forked from huawei/mindspore2022
dataset strategy set
This commit is contained in:
parent
ff5edd69dd
commit
dc7dc7d3fa
|
|
@ -73,6 +73,7 @@ void ParallelContext::Reset() {
|
|||
optimizer_weight_shard_aggregated_save_ = false;
|
||||
sharding_propagation_ = false;
|
||||
enable_all2all_ = false;
|
||||
dataset_strategy_.clear();
|
||||
}
|
||||
|
||||
void ParallelContext::set_device_num(int64_t device_num) {
|
||||
|
|
@ -89,6 +90,10 @@ void ParallelContext::set_gradients_mean(bool gradients_mean) { gradients_mean_
|
|||
|
||||
void ParallelContext::set_full_batch(bool full_batch) { full_batch_ = full_batch; }
|
||||
|
||||
void ParallelContext::set_dataset_strategy(const std::vector<std::vector<int64_t>> &dataset_strategy) {
|
||||
dataset_strategy_ = dataset_strategy;
|
||||
}
|
||||
|
||||
void ParallelContext::set_grad_accumulation_step(int64_t grad_accumulation_step) {
|
||||
grad_accumulation_step_ = grad_accumulation_step;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -65,6 +65,9 @@ class ParallelContext {
|
|||
void set_full_batch(bool full_batch);
|
||||
bool full_batch() const { return full_batch_; }
|
||||
|
||||
void set_dataset_strategy(const std::vector<std::vector<int64_t>> &dataset_strategy);
|
||||
std::vector<std::vector<int64_t>> dataset_strategy() const { return dataset_strategy_; }
|
||||
|
||||
void set_gradient_fp32_sync(bool gradient_fp32_sync);
|
||||
bool gradient_fp32_sync() const { return gradient_fp32_sync_; }
|
||||
|
||||
|
|
@ -169,6 +172,7 @@ class ParallelContext {
|
|||
bool sharding_propagation_;
|
||||
// Enable AllToAll or not. If false, use AllGather and Split.
|
||||
bool enable_all2all_;
|
||||
std::vector<std::vector<int64_t>> dataset_strategy_;
|
||||
};
|
||||
|
||||
} // namespace parallel
|
||||
|
|
|
|||
|
|
@ -30,19 +30,25 @@
|
|||
namespace mindspore {
|
||||
namespace parallel {
|
||||
Status GetNextInfo::InferTensorMap() {
|
||||
MS_EXCEPTION_IF_NULL(ParallelContext::GetInstance());
|
||||
bool full_batch = ParallelContext::GetInstance()->full_batch();
|
||||
|
||||
for (auto shp : shapes_) {
|
||||
TensorMap out_tensor_map;
|
||||
for (size_t i = 0; i < shp.size(); ++i) {
|
||||
if (full_batch) {
|
||||
out_tensor_map.push_back(MAP_NONE);
|
||||
auto slice_dim_iter = std::find(dev_matrix_shape_.begin(), dev_matrix_shape_.end(), shard_num_);
|
||||
if (slice_dim_iter == dev_matrix_shape_.end()) {
|
||||
MS_LOG(ERROR) << name_ << ": The dataset shard strategy only support shard in one dim.";
|
||||
return FAILED;
|
||||
}
|
||||
size_t slice_dim = size_t(slice_dim_iter - dev_matrix_shape_.begin());
|
||||
for (size_t i = 0; i < dataset_strategy_.size(); i++) {
|
||||
Shape tensor_map_index;
|
||||
for (auto dim : dataset_strategy_[i]) {
|
||||
if (dim == 1) {
|
||||
tensor_map_index.push_back(MAP_NONE);
|
||||
} else if (dim == shard_num_) {
|
||||
tensor_map_index.push_back(dev_matrix_shape_.size() - 1 - slice_dim);
|
||||
} else {
|
||||
out_tensor_map.push_back(SizeToLong(dev_matrix_shape_.size() - i - 1));
|
||||
MS_LOG(ERROR) << name_ << ": The dataset shard strategy only support fully shard in one dim.";
|
||||
return FAILED;
|
||||
}
|
||||
}
|
||||
outputs_tensor_map_.push_back(out_tensor_map);
|
||||
outputs_tensor_map_.push_back(tensor_map_index);
|
||||
}
|
||||
return SUCCESS;
|
||||
}
|
||||
|
|
@ -62,19 +68,6 @@ Status GetNextInfo::InferTensorLayout(TensorLayouts *outputs_layout) {
|
|||
return SUCCESS;
|
||||
}
|
||||
|
||||
Strategys GetNextInfo::GetOutputStrategy() {
|
||||
Strategys outputs_strategy;
|
||||
for (auto shp : shapes_) {
|
||||
Dimensions out_strategy;
|
||||
out_strategy.push_back(stage_device_size_);
|
||||
for (size_t i = 1; i < shp.size(); ++i) {
|
||||
out_strategy.push_back(1);
|
||||
}
|
||||
outputs_strategy.push_back(out_strategy);
|
||||
}
|
||||
return outputs_strategy;
|
||||
}
|
||||
|
||||
Status GetNextInfo::InferTensorInfo() {
|
||||
TensorLayouts outputs_layout;
|
||||
if (InferTensorLayout(&outputs_layout) != SUCCESS) {
|
||||
|
|
@ -88,23 +81,25 @@ Status GetNextInfo::InferTensorInfo() {
|
|||
}
|
||||
|
||||
Status GetNextInfo::InferDevMatrixShape() {
|
||||
size_t max_shape_length = 0;
|
||||
for (auto shp : shapes_) {
|
||||
if (max_shape_length < shp.size()) {
|
||||
max_shape_length = shp.size();
|
||||
}
|
||||
if (dataset_strategy_.empty()) {
|
||||
MS_LOG(ERROR) << "The dataset strategy is empty";
|
||||
return FAILED;
|
||||
}
|
||||
if (max_shape_length == 0) {
|
||||
MS_LOG(ERROR) << name_ << " : shape is 0";
|
||||
auto dev_matrix_iter =
|
||||
std::max_element(dataset_strategy_.begin(), dataset_strategy_.end(),
|
||||
[](Dimensions stra1, Dimensions stra2) { return stra1.size() < stra2.size(); });
|
||||
if (dev_matrix_iter != dataset_strategy_.end()) {
|
||||
dev_matrix_shape_ = *dev_matrix_iter;
|
||||
}
|
||||
dev_matrix_shape_.push_back(stage_device_size_);
|
||||
for (size_t i = 1; i < max_shape_length; ++i) {
|
||||
dev_matrix_shape_.push_back(1);
|
||||
auto shard_num_iter = std::max_element(dev_matrix_shape_.begin(), dev_matrix_shape_.end());
|
||||
if (shard_num_iter != dev_matrix_shape_.end()) {
|
||||
shard_num_ = *shard_num_iter;
|
||||
}
|
||||
return SUCCESS;
|
||||
}
|
||||
|
||||
Status GetNextInfo::Init(const StrategyPtr &strategy) {
|
||||
repeated_num_in_dev_matrix_right_ = false;
|
||||
if (InitWithAutoRepeatCalc(strategy) != SUCCESS) {
|
||||
MS_LOG(ERROR) << name_ << " : Init failed";
|
||||
return FAILED;
|
||||
|
|
@ -125,6 +120,25 @@ Status GetNextInfo::CheckStrategy(const StrategyPtr &strategy) {
|
|||
return FAILED;
|
||||
}
|
||||
}
|
||||
MS_EXCEPTION_IF_NULL(ParallelContext::GetInstance());
|
||||
if (!ParallelContext::GetInstance()->dataset_strategy().empty()) {
|
||||
dataset_strategy_ = ParallelContext::GetInstance()->dataset_strategy();
|
||||
} else {
|
||||
bool full_batch = ParallelContext::GetInstance()->full_batch();
|
||||
int64_t dev_num = full_batch ? 1 : SizeToLong(g_device_manager->stage_device_num());
|
||||
for (size_t i = 0; i < outputs_shape_.size(); i++) {
|
||||
Dimensions input_strategy;
|
||||
for (size_t j = 0; j < outputs_shape_[i].size(); j++) {
|
||||
input_strategy.push_back(1);
|
||||
}
|
||||
dataset_strategy_.push_back(input_strategy);
|
||||
}
|
||||
for (auto &stra : dataset_strategy_) {
|
||||
if (!stra.empty()) {
|
||||
stra[0] = dev_num;
|
||||
}
|
||||
}
|
||||
}
|
||||
return SUCCESS;
|
||||
}
|
||||
|
||||
|
|
@ -191,23 +205,9 @@ Status GetNextInfo::GetAttrs() {
|
|||
}
|
||||
|
||||
Status GetNextInfo::InferReplaceOps(const StrategyPtr &) {
|
||||
MS_EXCEPTION_IF_NULL(ParallelContext::GetInstance());
|
||||
bool full_batch = ParallelContext::GetInstance()->full_batch();
|
||||
|
||||
Shapes out_shapes = outputs_shape_;
|
||||
for (size_t i = 0; i < out_shapes.size(); ++i) {
|
||||
if (stage_device_size_ <= 0) {
|
||||
MS_LOG(ERROR) << name_ << " : The dev num is 0.";
|
||||
return FAILED;
|
||||
}
|
||||
if (!full_batch) {
|
||||
if (out_shapes[i][0] % stage_device_size_ != 0) {
|
||||
MS_LOG(ERROR) << name_ << " : batch num cannot floor div dev num.";
|
||||
return FAILED;
|
||||
}
|
||||
out_shapes[i][0] = out_shapes[i][0] / stage_device_size_;
|
||||
}
|
||||
}
|
||||
Shapes out_shapes;
|
||||
std::transform(outputs_tensor_info_.begin(), outputs_tensor_info_.end(), std::back_inserter(out_shapes),
|
||||
[](auto tensor_info) { return tensor_info.slice_shape(); });
|
||||
ValuePtr new_shapes = MakeValue(out_shapes);
|
||||
Attr attr_types = std::make_pair(TYPES, attrs_[TYPES]);
|
||||
Attr attr_shapes = std::make_pair(SHAPES, new_shapes);
|
||||
|
|
|
|||
|
|
@ -53,7 +53,6 @@ class GetNextInfo : public OperatorInfo {
|
|||
Status GetAttrTypes();
|
||||
Status GetAttrShapes();
|
||||
Status GetAttrOutPutNum();
|
||||
Strategys GetOutputStrategy();
|
||||
Status InferAsLossDivisor() override { return SUCCESS; }
|
||||
|
||||
private:
|
||||
|
|
@ -61,7 +60,9 @@ class GetNextInfo : public OperatorInfo {
|
|||
std::vector<std::string> types_;
|
||||
Shapes shapes_;
|
||||
int64_t output_num_ = 0;
|
||||
int64_t shard_num_ = 1;
|
||||
std::string shared_name_;
|
||||
Strategys dataset_strategy_;
|
||||
};
|
||||
} // namespace parallel
|
||||
} // namespace mindspore
|
||||
|
|
|
|||
|
|
@ -100,6 +100,8 @@ constexpr char STEP_AUTO_PARALLEL_BEGIN[] = "step_auto_parallel_begin.dot";
|
|||
constexpr char REQUIRES_GRAD[] = "requires_grad";
|
||||
constexpr char PARAM_NAME[] = "name";
|
||||
constexpr char RESHAPEINFO[] = "ReshapeInfo";
|
||||
constexpr char GETNEXTINFO[] = "GetNextInfo";
|
||||
constexpr char VIRTUALDATASETINFO[] = "VirtualDatasetInfo";
|
||||
|
||||
constexpr char RELU_TYPE[] = "relu";
|
||||
constexpr char RELU6_TYPE[] = "relu6";
|
||||
|
|
|
|||
|
|
@ -19,6 +19,8 @@
|
|||
#include <memory>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
#include <algorithm>
|
||||
#include <functional>
|
||||
|
||||
#include "frontend/parallel/device_manager.h"
|
||||
#include "frontend/parallel/device_matrix.h"
|
||||
|
|
@ -39,42 +41,51 @@ Status VirtualDatasetInfo::CheckStrategy(const StrategyPtr &strategy) {
|
|||
MS_LOG(ERROR) << name_ << ": Strategy size must be larger than 1.";
|
||||
return FAILED;
|
||||
}
|
||||
if (stra.size() == 1) {
|
||||
MS_LOG(WARNING) << name_ << ": Strategy size is 1.";
|
||||
return SUCCESS;
|
||||
}
|
||||
Dimensions strategy_first = stra.at(1);
|
||||
for (auto iter_strategy = stra.begin() + 1; iter_strategy != stra.end(); ++iter_strategy) {
|
||||
if (iter_strategy->empty()) {
|
||||
MS_LOG(ERROR) << name_ << ": iter_strategy size is zero.";
|
||||
}
|
||||
if (strategy_first.at(0) != *(iter_strategy->begin())) {
|
||||
MS_LOG(ERROR) << name_ << ": The first dimension of each strategy must be the same.";
|
||||
return FAILED;
|
||||
}
|
||||
|
||||
for (auto iter_element = iter_strategy->begin() + 1; iter_element != iter_strategy->end(); ++iter_element) {
|
||||
if (*iter_element != 1) {
|
||||
MS_LOG(ERROR) << name_ << ": All dimension except the first dimension of each strategy must be 1.";
|
||||
used_devices_ = int64_t(std::accumulate(stra[0].begin(), stra[0].end(), 1, std::multiplies<int64_t>()));
|
||||
for (size_t i = 0; i < stra.size(); ++i) {
|
||||
bool find_shard_dim = false;
|
||||
int64_t current_stra_shard_num = 1;
|
||||
for (auto dim : stra[i]) {
|
||||
if (dim == 1) {
|
||||
continue;
|
||||
}
|
||||
if (find_shard_dim) {
|
||||
MS_LOG(ERROR) << name_ << ": The dataset shard strategy only support shard in one dim.";
|
||||
return FAILED;
|
||||
} else {
|
||||
find_shard_dim = true;
|
||||
current_stra_shard_num = dim;
|
||||
}
|
||||
}
|
||||
if (i == 0) {
|
||||
shard_num_ = current_stra_shard_num;
|
||||
} else if (current_stra_shard_num != 1 && current_stra_shard_num != shard_num_) {
|
||||
MS_LOG(ERROR) << name_
|
||||
<< ": For each dataset input, the shard strategy can be not shard, "
|
||||
"or shard in one dim with the same shard size between each input. "
|
||||
"Current shard size is: "
|
||||
<< current_stra_shard_num << ". The previous shard size is " << shard_num_;
|
||||
return FAILED;
|
||||
}
|
||||
if (stra[i].size() > stra[max_size_strategy_dim_].size()) {
|
||||
max_size_strategy_dim_ = i;
|
||||
}
|
||||
}
|
||||
if (std::find(stra[max_size_strategy_dim_].begin(), stra[max_size_strategy_dim_].end(), shard_num_) ==
|
||||
stra[max_size_strategy_dim_].end()) {
|
||||
MS_LOG(ERROR) << name_
|
||||
<< ": For each dataset input, the shard strategy can be not shard, "
|
||||
"or shard in one dim with the same shard size between each input."
|
||||
" If using shard, the max length input must be shard, "
|
||||
"but the strategy of the max length input is: "
|
||||
<< stra[max_size_strategy_dim_];
|
||||
}
|
||||
return SUCCESS;
|
||||
}
|
||||
|
||||
Status VirtualDatasetInfo::InferDevMatrixShape() {
|
||||
Strategys stra = strategy_->GetInputDim();
|
||||
Dimensions strategy_first = stra.at(0);
|
||||
int64_t batch_split_num = ((int64_t)(strategy_first.at(0)));
|
||||
dev_matrix_shape_.push_back(batch_split_num);
|
||||
if (stage_device_size_ > batch_split_num) {
|
||||
dev_matrix_shape_.push_back(stage_device_size_ / batch_split_num);
|
||||
}
|
||||
// Because 'VirtualDataSet' uses 'InitWithManualRepeatCalc' which does not calculates 'used_devices_',
|
||||
// we calculate it here.
|
||||
used_devices_ = batch_split_num;
|
||||
|
||||
dev_matrix_shape_ = stra[max_size_strategy_dim_];
|
||||
return SUCCESS;
|
||||
}
|
||||
|
||||
|
|
@ -83,18 +94,24 @@ Status VirtualDatasetInfo::InferMirrorOps() { return SUCCESS; }
|
|||
Status VirtualDatasetInfo::InferForwardCommunication() { return SUCCESS; }
|
||||
|
||||
Status VirtualDatasetInfo::InferTensorMap() {
|
||||
MS_EXCEPTION_IF_NULL(ParallelContext::GetInstance());
|
||||
bool full_batch = ParallelContext::GetInstance()->full_batch();
|
||||
|
||||
for (size_t i = 0; i < strategy_->GetInputNumber(); i++) {
|
||||
auto slice_dim_iter = std::find(dev_matrix_shape_.begin(), dev_matrix_shape_.end(), shard_num_);
|
||||
if (slice_dim_iter == dev_matrix_shape_.end()) {
|
||||
MS_LOG(ERROR) << name_ << ": The dataset shard strategy only support shard in one dim.";
|
||||
return FAILED;
|
||||
}
|
||||
size_t slice_dim = size_t(slice_dim_iter - dev_matrix_shape_.begin());
|
||||
auto stra = strategy_->GetInputDim();
|
||||
for (size_t i = 0; i < stra.size(); i++) {
|
||||
Shape tensor_map_index;
|
||||
if (full_batch) {
|
||||
tensor_map_index.push_back(MAP_NONE);
|
||||
} else {
|
||||
tensor_map_index.push_back((int64_t)(LAST_INDEX(dev_matrix_shape_.size())));
|
||||
}
|
||||
for (size_t j = 1; j < strategy_->GetInputDim()[i].size(); ++j) {
|
||||
tensor_map_index.push_back(MAP_NONE);
|
||||
for (auto dim : stra[i]) {
|
||||
if (dim == 1) {
|
||||
tensor_map_index.push_back(MAP_NONE);
|
||||
} else if (dim == shard_num_) {
|
||||
tensor_map_index.push_back(dev_matrix_shape_.size() - 1 - slice_dim);
|
||||
} else {
|
||||
MS_LOG(ERROR) << name_ << ": The dataset shard strategy only support shard in one dim.";
|
||||
return FAILED;
|
||||
}
|
||||
}
|
||||
inputs_tensor_map_.push_back(tensor_map_index);
|
||||
outputs_tensor_map_.push_back(tensor_map_index);
|
||||
|
|
@ -105,7 +122,8 @@ Status VirtualDatasetInfo::InferTensorMap() {
|
|||
Status VirtualDatasetInfo::GetAttrs() { return SUCCESS; }
|
||||
|
||||
Status VirtualDatasetInfo::Init(const StrategyPtr &strategy) {
|
||||
if (InitWithManualRepeatCalc(strategy) != SUCCESS) {
|
||||
repeated_num_in_dev_matrix_right_ = false;
|
||||
if (InitWithAutoRepeatCalc(strategy) != SUCCESS) {
|
||||
MS_LOG(ERROR) << name_ << ": Init failed.";
|
||||
return FAILED;
|
||||
}
|
||||
|
|
@ -113,7 +131,7 @@ Status VirtualDatasetInfo::Init(const StrategyPtr &strategy) {
|
|||
}
|
||||
|
||||
Status VirtualDatasetInfo::InitForCostModel(const StrategyPtr &strategy) {
|
||||
if (InitForCostModelWithManualRepeatCalc(strategy) != SUCCESS) {
|
||||
if (InitForCostModelWithAutoRepeatCalc(strategy) != SUCCESS) {
|
||||
MS_LOG(ERROR) << name_ << ": Init for cost model failed.";
|
||||
return FAILED;
|
||||
}
|
||||
|
|
@ -134,42 +152,31 @@ Status VirtualDatasetInfo::SetCostUnderStrategy(const StrategyPtr &strategy) {
|
|||
|
||||
Status VirtualDatasetInfo::GenerateStrategies(int64_t stage_id) {
|
||||
MS_EXCEPTION_IF_NULL(ParallelContext::GetInstance());
|
||||
bool full_batch = ParallelContext::GetInstance()->full_batch();
|
||||
size_t total_dev_num;
|
||||
|
||||
if (GetAttrs() != SUCCESS) {
|
||||
MS_LOG(ERROR) << name_ << ": GetAttrs failed";
|
||||
return FAILED;
|
||||
}
|
||||
|
||||
if (full_batch) {
|
||||
total_dev_num = 1;
|
||||
} else {
|
||||
total_dev_num = stage_device_size_;
|
||||
}
|
||||
StrategyPtr sp;
|
||||
Strategys strategy;
|
||||
for (auto &shape : inputs_shape_) {
|
||||
Shape temp;
|
||||
temp.emplace_back(SizeToLong(total_dev_num));
|
||||
(void)temp.insert(temp.end(), shape.size() - 1, 1);
|
||||
strategy.push_back(temp);
|
||||
if (!ParallelContext::GetInstance()->dataset_strategy().empty()) {
|
||||
strategy = ParallelContext::GetInstance()->dataset_strategy();
|
||||
} else {
|
||||
bool full_batch = ParallelContext::GetInstance()->full_batch();
|
||||
size_t total_dev_num;
|
||||
if (full_batch) {
|
||||
total_dev_num = 1;
|
||||
} else {
|
||||
total_dev_num = stage_device_size_;
|
||||
}
|
||||
for (auto &shape : inputs_shape_) {
|
||||
Shape temp;
|
||||
temp.emplace_back(SizeToLong(total_dev_num));
|
||||
(void)temp.insert(temp.end(), shape.size() - 1, 1);
|
||||
strategy.push_back(temp);
|
||||
}
|
||||
}
|
||||
sp = std::make_shared<Strategy>(stage_id, strategy);
|
||||
|
||||
if (SetCostUnderStrategy(sp) == SUCCESS) {
|
||||
if (full_batch) {
|
||||
MS_LOG(INFO) << name_ << ": Successfully generated full-batch-parallel-strategy.";
|
||||
} else {
|
||||
MS_LOG(INFO) << name_ << ": Successfully generated batch-parallel-strategy.";
|
||||
}
|
||||
MS_LOG(INFO) << name_ << ": Successfully dataset strategy.";
|
||||
PrintStrategy(sp);
|
||||
} else {
|
||||
if (full_batch) {
|
||||
MS_LOG(ERROR) << name_ << ": Generating full-batch-parallel-strategy failed.";
|
||||
} else {
|
||||
MS_LOG(ERROR) << name_ << ": Generating batch-parallel-strategy failed.";
|
||||
}
|
||||
MS_LOG(ERROR) << name_ << ": Generating dataset strategy failed.";
|
||||
return FAILED;
|
||||
}
|
||||
return SUCCESS;
|
||||
|
|
|
|||
|
|
@ -50,6 +50,8 @@ class VirtualDatasetInfo : public OperatorInfo {
|
|||
Status InferTensorMap() override;
|
||||
Status GetAttrs() override;
|
||||
Status InferAsLossDivisor() override;
|
||||
size_t max_size_strategy_dim_ = 0;
|
||||
int64_t shard_num_ = 1;
|
||||
};
|
||||
} // namespace parallel
|
||||
} // namespace mindspore
|
||||
|
|
|
|||
|
|
@ -46,8 +46,38 @@ Status VirtualOutputInfo::CheckStrategy(const StrategyPtr &strategy) {
|
|||
return FAILED;
|
||||
}
|
||||
}
|
||||
|
||||
if (!strategy_first.empty()) {
|
||||
shard_num_ = strategy_first[0];
|
||||
}
|
||||
return SUCCESS;
|
||||
}
|
||||
|
||||
Status VirtualOutputInfo::GenerateStrategies(int64_t stage_id) {
|
||||
StrategyPtr sp;
|
||||
Strategys strategy;
|
||||
bool full_batch = ParallelContext::GetInstance()->full_batch();
|
||||
size_t total_dev_num;
|
||||
if (full_batch) {
|
||||
total_dev_num = 1;
|
||||
} else {
|
||||
total_dev_num = stage_device_size_;
|
||||
}
|
||||
for (auto &shape : inputs_shape_) {
|
||||
Shape temp;
|
||||
temp.emplace_back(SizeToLong(total_dev_num));
|
||||
(void)temp.insert(temp.end(), shape.size() - 1, 1);
|
||||
strategy.push_back(temp);
|
||||
}
|
||||
sp = std::make_shared<Strategy>(stage_id, strategy);
|
||||
if (SetCostUnderStrategy(sp) == SUCCESS) {
|
||||
MS_LOG(INFO) << name_ << ": Successfully dataset strategy.";
|
||||
PrintStrategy(sp);
|
||||
} else {
|
||||
MS_LOG(ERROR) << name_ << ": Generating dataset strategy failed.";
|
||||
return FAILED;
|
||||
}
|
||||
return SUCCESS;
|
||||
}
|
||||
|
||||
} // namespace parallel
|
||||
} // namespace mindspore
|
||||
|
|
|
|||
|
|
@ -35,6 +35,7 @@ class VirtualOutputInfo : public VirtualDatasetInfo {
|
|||
const PrimitiveAttrs &attrs)
|
||||
: VirtualDatasetInfo(name, inputs_shape, outputs_shape, attrs) {}
|
||||
~VirtualOutputInfo() override = default;
|
||||
Status GenerateStrategies(int64_t stage_id) override;
|
||||
|
||||
protected:
|
||||
Status CheckStrategy(const StrategyPtr &strategy) override;
|
||||
|
|
|
|||
|
|
@ -95,7 +95,7 @@ bool StepAutoParallel(const FuncGraphPtr &root, const opt::OptimizerPtr &) {
|
|||
}
|
||||
// mark the forward cnodes, parallel only care these nodes
|
||||
MarkForwardCNode(root);
|
||||
if (!root->has_flag(TRAINING)) {
|
||||
if (IsInsertVirtualOutput(root)) {
|
||||
InsertVirtualOutput(root, all_nodes);
|
||||
AnfNodePtr ret_after = root->get_return();
|
||||
MS_EXCEPTION_IF_NULL(ret_after);
|
||||
|
|
|
|||
|
|
@ -2012,13 +2012,23 @@ void SetVirtualDatasetStrategy(const CNodePtr &node) {
|
|||
MS_EXCEPTION_IF_NULL(prim);
|
||||
if (prim->name() == VIRTUAL_DATA_SET || prim->name() == VIRTUAL_OUTPUT) {
|
||||
CheckGlobalDeviceManager();
|
||||
auto attrs_temp = prim->attrs();
|
||||
if (!ParallelContext::GetInstance()->dataset_strategy().empty() && prim->name() == VIRTUAL_DATA_SET) {
|
||||
std::vector<ValuePtr> elements;
|
||||
auto dataset_strategy = ParallelContext::GetInstance()->dataset_strategy();
|
||||
std::transform(dataset_strategy.begin(), dataset_strategy.end(), std::back_inserter(elements),
|
||||
[](auto input_stra) { return MakeValue(input_stra); });
|
||||
ValueTuplePtr strategy = std::make_shared<ValueTuple>(elements);
|
||||
attrs_temp[STRATEGY] = strategy;
|
||||
(void)prim->SetAttrs(attrs_temp);
|
||||
return;
|
||||
}
|
||||
int64_t dev_num;
|
||||
if (full_batch) {
|
||||
dev_num = 1;
|
||||
} else {
|
||||
dev_num = SizeToLong(g_device_manager->stage_device_num());
|
||||
}
|
||||
auto attrs_temp = prim->attrs();
|
||||
std::vector<Shapes> shape_list = ExtractShape(node);
|
||||
if (shape_list.empty()) {
|
||||
MS_LOG(EXCEPTION) << "Failure:node " << node->ToString() << " failed to extract shape";
|
||||
|
|
@ -3702,6 +3712,11 @@ void ReorderForPipelineSplit(const FuncGraphPtr &root, const FuncGraphManagerPtr
|
|||
}
|
||||
}
|
||||
|
||||
bool IsInsertVirtualOutput(const FuncGraphPtr &root) {
|
||||
MS_EXCEPTION_IF_NULL(ParallelContext::GetInstance());
|
||||
return (!root->has_flag(TRAINING) && ParallelContext::GetInstance()->dataset_strategy().empty());
|
||||
}
|
||||
|
||||
bool StepParallel(const FuncGraphPtr &root, const opt::OptimizerPtr &optimizer) {
|
||||
#if (ENABLE_CPU && !_WIN32)
|
||||
if (ps::PSContext::instance()->is_server() || ps::PSContext::instance()->is_scheduler()) {
|
||||
|
|
@ -3762,7 +3777,7 @@ bool StepParallel(const FuncGraphPtr &root, const opt::OptimizerPtr &optimizer)
|
|||
MS_LOG(EXCEPTION) << "The graph contain communication op";
|
||||
}
|
||||
|
||||
if (!root->has_flag(TRAINING)) {
|
||||
if (IsInsertVirtualOutput(root)) {
|
||||
InsertVirtualOutput(root, all_nodes);
|
||||
AnfNodePtr ret_after = root->get_return();
|
||||
MS_EXCEPTION_IF_NULL(ret_after);
|
||||
|
|
|
|||
|
|
@ -116,7 +116,7 @@ std::string SetParallelShape(const AnfNodePtr ¶meter, const std::pair<AnfNod
|
|||
void CoverSliceShape(const FuncGraphPtr &root);
|
||||
|
||||
void SetVirtualDatasetStrategy(const CNodePtr &node);
|
||||
|
||||
bool IsInsertVirtualOutput(const FuncGraphPtr &root);
|
||||
// Create parallel operator for primitive node(has strategy)
|
||||
void ExtractInformation(const std::vector<AnfNodePtr> &all_nodes, bool is_training = true);
|
||||
|
||||
|
|
|
|||
|
|
@ -170,6 +170,8 @@ PYBIND11_MODULE(_c_expression, m) {
|
|||
.def("get_pipeline_stage_split_num", &ParallelContext::pipeline_stage_split_num, "Get pipeline stage split num.")
|
||||
.def("set_full_batch", &ParallelContext::set_full_batch, "Set whether load full batch on each device.")
|
||||
.def("get_full_batch", &ParallelContext::full_batch, "Get whether load full batch on each device.")
|
||||
.def("set_dataset_strategy", &ParallelContext::set_dataset_strategy, "Set dataset sharding strategy.")
|
||||
.def("get_dataset_strategy", &ParallelContext::dataset_strategy, "Get dataset sharding strategy.")
|
||||
.def("set_enable_parallel_optimizer", &ParallelContext::set_enable_parallel_optimizer,
|
||||
"Set enable/disable parallel optimizer.")
|
||||
.def("get_enable_parallel_optimizer", &ParallelContext::enable_parallel_optimizer,
|
||||
|
|
|
|||
|
|
@ -256,6 +256,31 @@ class _AutoParallelContext:
|
|||
return False
|
||||
return self._context_handle.get_full_batch()
|
||||
|
||||
def set_dataset_strategy(self, dataset_strategy):
|
||||
"""
|
||||
Set dataset sharding strategy.
|
||||
|
||||
Args:
|
||||
dataset_strategy (tuple(tuple)): The dataset sharding strategy.
|
||||
"""
|
||||
self.check_context_handle()
|
||||
if not isinstance(dataset_strategy, tuple):
|
||||
raise TypeError(f'strategy must be tuple type, but got:{type(dataset_strategy)}')
|
||||
for ele in dataset_strategy:
|
||||
if not isinstance(ele, tuple):
|
||||
raise TypeError(f'The element of strategy must be tuple type, but got:{type(ele)}')
|
||||
for dim in ele:
|
||||
if not isinstance(dim, int):
|
||||
raise TypeError(f'The dim of each strategy value must be int type, but got:{type(dim)}')
|
||||
self._context_handle.set_dataset_strategy(dataset_strategy)
|
||||
|
||||
def get_dataset_strategy(self):
|
||||
"""Get dataset sharding strategy."""
|
||||
self.check_context_handle()
|
||||
if _is_role_pserver():
|
||||
return False
|
||||
return self._context_handle.get_dataset_strategy()
|
||||
|
||||
def set_grad_accumulation_step(self, grad_accumulation_step):
|
||||
"""
|
||||
Set grad accumulation step.
|
||||
|
|
@ -596,6 +621,7 @@ _set_auto_parallel_context_func_map = {
|
|||
"strategy_ckpt_save_file": auto_parallel_context().set_strategy_ckpt_save_file,
|
||||
"group_ckpt_save_file": auto_parallel_context().set_group_ckpt_save_file,
|
||||
"full_batch": auto_parallel_context().set_full_batch,
|
||||
"dataset_strategy": auto_parallel_context().set_dataset_strategy,
|
||||
"enable_parallel_optimizer": auto_parallel_context().set_enable_parallel_optimizer,
|
||||
"grad_accumulation_step": auto_parallel_context().set_grad_accumulation_step,
|
||||
"all_reduce_fusion_config": auto_parallel_context().set_all_reduce_fusion_split_indices,
|
||||
|
|
@ -619,6 +645,7 @@ _get_auto_parallel_context_func_map = {
|
|||
"strategy_ckpt_load_file": auto_parallel_context().get_strategy_ckpt_load_file,
|
||||
"strategy_ckpt_save_file": auto_parallel_context().get_strategy_ckpt_save_file,
|
||||
"full_batch": auto_parallel_context().get_full_batch,
|
||||
"dataset_strategy": auto_parallel_context().get_dataset_strategy,
|
||||
"enable_parallel_optimizer": auto_parallel_context().get_enable_parallel_optimizer,
|
||||
"grad_accumulation_step": auto_parallel_context().get_grad_accumulation_step,
|
||||
"all_reduce_fusion_config": auto_parallel_context().get_all_reduce_fusion_split_indices,
|
||||
|
|
@ -632,7 +659,7 @@ _get_auto_parallel_context_func_map = {
|
|||
@args_type_check(device_num=int, global_rank=int, gradients_mean=bool, gradient_fp32_sync=bool,
|
||||
loss_repeated_mean=bool, parallel_mode=str, auto_parallel_search_mode=str,
|
||||
parameter_broadcast=bool, strategy_ckpt_load_file=str,
|
||||
strategy_ckpt_save_file=str, full_batch=bool, enable_parallel_optimizer=bool,
|
||||
strategy_ckpt_save_file=str, full_batch=bool, dataset_strategy=tuple, enable_parallel_optimizer=bool,
|
||||
grad_accumulation_step=int, all_reduce_fusion_config=list, group_ckpt_save_file=str,
|
||||
communi_parallel_mode=str, optimizer_weight_shard_size=int,
|
||||
optimizer_weight_shard_aggregated_save=bool,
|
||||
|
|
@ -679,6 +706,7 @@ def _set_auto_parallel_context(**kwargs):
|
|||
strategy_ckpt_save_file (str): The path to save parallel strategy checkpoint. Default: ''
|
||||
group_ckpt_save_file (str): The path to save parallel group checkpoint. Default: ''
|
||||
full_batch (bool): Whether to load the whole batch on each device. Default: False.
|
||||
dataset_strategy (tuplr): Dataset sharding strategy. Default: ().
|
||||
enable_parallel_optimizer (bool): Enable using optimizer segmentation or not. Default: False.
|
||||
all_reduce_fusion_config (list): Set allreduce fusion strategy by parameters indices.
|
||||
pipeline_stages (int): Set the stage information for pipeline parallel. This indicates how
|
||||
|
|
|
|||
|
|
@ -58,6 +58,9 @@ def _check_full_batch():
|
|||
|
||||
def _need_to_full():
|
||||
"""Check whether to convert input to full shape or tensor."""
|
||||
dataset_strategy = context.get_auto_parallel_context("dataset_strategy")
|
||||
if dataset_strategy:
|
||||
return True
|
||||
parallel_mode = _get_parallel_mode()
|
||||
full_batch = _get_full_batch()
|
||||
need = ((parallel_mode in ("semi_auto_parallel", "auto_parallel"))
|
||||
|
|
@ -68,6 +71,20 @@ def _need_to_full():
|
|||
def _to_full_shapes(shapes, device_num):
|
||||
"""Expanding batch dimension according to device_num, adapt to mindspore minddata graph solution."""
|
||||
new_shapes = []
|
||||
dataset_strategy = context.get_auto_parallel_context("dataset_strategy")
|
||||
if dataset_strategy:
|
||||
if len(shapes) != len(dataset_strategy):
|
||||
raise ValueError("The input shapes size {} is not equal to "
|
||||
"dataset strategy size {}".format(len(shapes), len(dataset_strategy)))
|
||||
for index, shape in enumerate(shapes):
|
||||
if len(shape) != len(dataset_strategy[index]):
|
||||
raise ValueError("The input shapes item size {} is not equal to "
|
||||
"dataset strategy item size {}".format(len(shape), len(dataset_strategy[index])))
|
||||
new_shape = ()
|
||||
for i, item in enumerate(shape):
|
||||
new_shape += (item * dataset_strategy[index][i],)
|
||||
new_shapes.append(new_shape)
|
||||
return new_shapes
|
||||
for shape in shapes:
|
||||
new_shape = ()
|
||||
for i, item in enumerate(shape):
|
||||
|
|
@ -91,8 +108,12 @@ def _to_full_tensor(elem, global_device_num, global_rank, scaling_sens=None):
|
|||
if stage_rank >= device_num:
|
||||
raise ValueError("The global rank must be smaller than device number, the global rank is {}, "
|
||||
"the device num is {}".format(stage_rank, device_num))
|
||||
|
||||
for data in elem:
|
||||
dataset_strategy = context.get_auto_parallel_context("dataset_strategy")
|
||||
if elem and dataset_strategy:
|
||||
if len(elem) != len(dataset_strategy):
|
||||
raise ValueError("The input size {} is not equal to "
|
||||
"dataset strategy size {}".format(len(elem), len(dataset_strategy)))
|
||||
for index, data in enumerate(elem):
|
||||
if isinstance(data, np.ndarray):
|
||||
data = Tensor(data)
|
||||
if not isinstance(data, Tensor):
|
||||
|
|
@ -100,16 +121,30 @@ def _to_full_tensor(elem, global_device_num, global_rank, scaling_sens=None):
|
|||
shape_ = data.shape
|
||||
type_ = data.dtype
|
||||
new_shape = ()
|
||||
batchsize_per_device = 1
|
||||
for i, item in enumerate(shape_):
|
||||
if i == 0:
|
||||
new_shape += (item * device_num,)
|
||||
batchsize_per_device = item
|
||||
else:
|
||||
new_shape += (item,)
|
||||
new_tensor_numpy = np.zeros(new_shape, dtype_to_nptype(type_))
|
||||
start = stage_rank * batchsize_per_device
|
||||
new_tensor_numpy[start: start + batchsize_per_device] = data.asnumpy()
|
||||
if not dataset_strategy:
|
||||
batchsize_per_device = 1
|
||||
for i, item in enumerate(shape_):
|
||||
if i == 0:
|
||||
new_shape += (item * device_num,)
|
||||
batchsize_per_device = item
|
||||
else:
|
||||
new_shape += (item,)
|
||||
new_tensor_numpy = np.zeros(new_shape, dtype_to_nptype(type_))
|
||||
start = stage_rank * batchsize_per_device
|
||||
new_tensor_numpy[start: start + batchsize_per_device] = data.asnumpy()
|
||||
else:
|
||||
if len(shape_) != len(dataset_strategy[index]):
|
||||
raise ValueError("The input shapes item size {} is not equal to "
|
||||
"dataset strategy item size {}".format(len(shape_), len(dataset_strategy[index])))
|
||||
slice_index = ()
|
||||
for i, item in enumerate(shape_):
|
||||
new_shape += (item * dataset_strategy[index][i],)
|
||||
start = (stage_rank % dataset_strategy[index][i]) * item
|
||||
end = (stage_rank % dataset_strategy[index][i] + 1) * item
|
||||
s = slice(start, end, 1)
|
||||
slice_index += (s,)
|
||||
new_tensor_numpy = np.zeros(new_shape, dtype_to_nptype(type_))
|
||||
new_tensor_numpy[slice_index] = data.asnumpy()
|
||||
new_tensor = Tensor(new_tensor_numpy)
|
||||
lst.append(new_tensor)
|
||||
if scaling_sens:
|
||||
|
|
|
|||
|
|
@ -67,86 +67,10 @@ TEST_F(TestVirtualDatasetInfo, InferDevMatrixShape1) {
|
|||
virtual_dataset->Init(strategy);
|
||||
Shape dev_matrix_shape = virtual_dataset->dev_matrix_shape();
|
||||
|
||||
Shape expect = {16};
|
||||
Shape expect = {16, 1};
|
||||
ASSERT_EQ(dev_matrix_shape, expect);
|
||||
}
|
||||
|
||||
TEST_F(TestVirtualDatasetInfo, InferDevMatrixShape2) {
|
||||
Strategys inputs = {{8, 1}, {8, 1}, {8, 1}};
|
||||
StrategyPtr strategy = NewStrategy(0, inputs);
|
||||
virtual_dataset->Init(strategy);
|
||||
Shape dev_matrix_shape = virtual_dataset->dev_matrix_shape();
|
||||
|
||||
Shape expect = {8, 2};
|
||||
ASSERT_EQ(dev_matrix_shape, expect);
|
||||
}
|
||||
|
||||
TEST_F(TestVirtualDatasetInfo, InferSliceShape1) {
|
||||
Strategys str = {{8, 1}, {8, 1}, {8, 1}};
|
||||
StrategyPtr strategy = NewStrategy(0, str);
|
||||
|
||||
virtual_dataset->Init(strategy);
|
||||
std::vector<TensorInfo> inputs = virtual_dataset->inputs_tensor_info();
|
||||
std::vector<TensorInfo> outputs = virtual_dataset->outputs_tensor_info();
|
||||
|
||||
Shape input_slice_shape_expect = {16, 32};
|
||||
Shape output_slice_shape_expect = {16, 32};
|
||||
|
||||
TensorInfo input_tensor_info = inputs.at(0);
|
||||
TensorInfo output_tensor_info = outputs.at(0);
|
||||
|
||||
Shape input_slice_shape = input_tensor_info.slice_shape();
|
||||
Shape output_slice_shape = output_tensor_info.slice_shape();
|
||||
|
||||
ASSERT_EQ(input_slice_shape, input_slice_shape_expect);
|
||||
ASSERT_EQ(output_slice_shape, output_slice_shape_expect);
|
||||
|
||||
Shape input_slice_shape_expect1 = {160, 320};
|
||||
Shape output_slice_shape_expect1 = {160, 320};
|
||||
|
||||
TensorInfo input_tensor_info1 = inputs.at(1);
|
||||
TensorInfo output_tensor_info1 = outputs.at(1);
|
||||
|
||||
Shape input_slice_shape1 = input_tensor_info1.slice_shape();
|
||||
Shape output_slice_shape1 = output_tensor_info1.slice_shape();
|
||||
|
||||
ASSERT_EQ(input_slice_shape1, input_slice_shape_expect1);
|
||||
ASSERT_EQ(output_slice_shape1, output_slice_shape_expect1);
|
||||
|
||||
Shape input_slice_shape_expect2 = {1600, 3200};
|
||||
Shape output_slice_shape_expect2 = {1600, 3200};
|
||||
|
||||
TensorInfo input_tensor_info2 = inputs.at(2);
|
||||
TensorInfo output_tensor_info2 = outputs.at(2);
|
||||
|
||||
Shape input_slice_shape2 = input_tensor_info2.slice_shape();
|
||||
Shape output_slice_shape2 = output_tensor_info2.slice_shape();
|
||||
|
||||
ASSERT_EQ(input_slice_shape2, input_slice_shape_expect2);
|
||||
ASSERT_EQ(output_slice_shape2, output_slice_shape_expect2);
|
||||
}
|
||||
|
||||
TEST_F(TestVirtualDatasetInfo, GetTensorLayout1) {
|
||||
Strategys str = {{8, 1}, {8, 1}, {8, 1}};
|
||||
StrategyPtr strategy = NewStrategy(0, str);
|
||||
|
||||
virtual_dataset->Init(strategy);
|
||||
std::vector<TensorInfo> inputs = virtual_dataset->inputs_tensor_info();
|
||||
std::vector<TensorInfo> outputs = virtual_dataset->outputs_tensor_info();
|
||||
|
||||
TensorMap input_expect = {1, -1};
|
||||
TensorMap output_expect = {1, -1};
|
||||
|
||||
TensorInfo input_tensor_info = inputs.at(0);
|
||||
TensorInfo output_tensor_info = outputs.at(0);
|
||||
|
||||
Map input_tensor_map = input_tensor_info.tensor_layout().origin_tensor_map();
|
||||
Map output_tensor_map = output_tensor_info.tensor_layout().origin_tensor_map();
|
||||
|
||||
ASSERT_EQ(input_tensor_map.array(), input_expect);
|
||||
ASSERT_EQ(output_tensor_map.array(), output_expect);
|
||||
}
|
||||
|
||||
TEST_F(TestVirtualDatasetInfo, GetForwardOp1) {
|
||||
Strategys inputs = {{8, 1}, {8, 1}, {8, 1}};
|
||||
StrategyPtr strategy = NewStrategy(0, inputs);
|
||||
|
|
|
|||
|
|
@ -22,7 +22,6 @@ from mindspore.common.api import _executor
|
|||
from mindspore.nn.wrap.cell_wrapper import VirtualDatasetCellTriple
|
||||
from mindspore.ops import composite as C
|
||||
from mindspore.ops import operations as P
|
||||
from mindspore.ops.operations.comm_ops import _VirtualDataset
|
||||
from tests.ut.python.ops.test_math_ops import VirtualLoss
|
||||
|
||||
|
||||
|
|
@ -48,38 +47,6 @@ class GradWrap(nn.Cell):
|
|||
def construct(self, x, y, b):
|
||||
return grad_all(self.network)(x, y, b)
|
||||
|
||||
|
||||
# model_parallel test
|
||||
def test_virtual_dataset_3_input():
|
||||
class Net(nn.Cell):
|
||||
def __init__(self, strategy0, strategy1, strategy2, strategy3):
|
||||
super().__init__()
|
||||
self.virtual_dataset = _VirtualDataset().shard(strategy0)
|
||||
self.matmul1 = P.MatMul().shard(strategy1)
|
||||
self.matmul2 = P.MatMul().shard(strategy2)
|
||||
self.gelu = P.GeLU().shard(strategy3)
|
||||
|
||||
def construct(self, x, y, b):
|
||||
x, y, b = self.virtual_dataset(x, y, b)
|
||||
out = self.gelu(self.matmul1(x, y))
|
||||
out = self.matmul2(out, b)
|
||||
return out
|
||||
|
||||
context.set_auto_parallel_context(parallel_mode="semi_auto_parallel")
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0)
|
||||
strategy0 = ((2, 1), (2, 1), (2, 1))
|
||||
strategy1 = ((2, 2), (2, 2))
|
||||
strategy2 = ((2, 2), (2, 2))
|
||||
strategy3 = ((2, 4),)
|
||||
net = GradWrap(NetWithLoss(Net(strategy0, strategy1, strategy2, strategy3)))
|
||||
x = Tensor(np.ones([128, 32]), dtype=ms.float32)
|
||||
y = Tensor(np.ones([32, 64]), dtype=ms.float32)
|
||||
b = Tensor(np.ones([64, 2048]), dtype=ms.float32)
|
||||
net.set_auto_parallel()
|
||||
net.set_train()
|
||||
_executor.compile(net, x, y, b)
|
||||
|
||||
|
||||
def test_virtualdataset_cell_3_inputs():
|
||||
class Net(nn.Cell):
|
||||
def __init__(self, strategy1, strategy2, strategy3):
|
||||
|
|
@ -106,5 +73,4 @@ def test_virtualdataset_cell_3_inputs():
|
|||
|
||||
|
||||
if __name__ == '__main__':
|
||||
test_virtual_dataset_3_input()
|
||||
context.reset_auto_parallel_context()
|
||||
|
|
|
|||
|
|
@ -0,0 +1,170 @@
|
|||
# Copyright 2021 Huawei Technologies Co., Ltd
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
import numpy as np
|
||||
|
||||
import mindspore as ms
|
||||
import mindspore.nn as nn
|
||||
from mindspore import Tensor
|
||||
from mindspore import context
|
||||
from mindspore.common.api import _executor
|
||||
from mindspore.ops import composite as C
|
||||
from mindspore.ops import operations as P
|
||||
from mindspore.ops.operations.comm_ops import _VirtualDataset
|
||||
from tests.ut.python.ops.test_math_ops import VirtualLoss
|
||||
grad_all = C.GradOperation(get_all=True)
|
||||
|
||||
class NetWithLoss(nn.Cell):
|
||||
def __init__(self, network):
|
||||
super(NetWithLoss, self).__init__()
|
||||
self.loss = VirtualLoss()
|
||||
self.network = network
|
||||
|
||||
def construct(self, x, y, b):
|
||||
predict = self.network(x, y, b)
|
||||
return self.loss(predict)
|
||||
|
||||
|
||||
class GradWrap(nn.Cell):
|
||||
def __init__(self, network):
|
||||
super(GradWrap, self).__init__()
|
||||
self.network = network
|
||||
|
||||
def construct(self, x, y, b):
|
||||
return grad_all(self.network)(x, y, b)
|
||||
|
||||
|
||||
class Net1(nn.Cell):
|
||||
def __init__(self, strategy1, strategy2, strategy3):
|
||||
super().__init__()
|
||||
self.virtual_dataset = _VirtualDataset()
|
||||
self.matmul1 = P.MatMul().shard(strategy1)
|
||||
self.matmul2 = P.MatMul().shard(strategy2)
|
||||
self.gelu = P.GeLU().shard(strategy3)
|
||||
|
||||
def construct(self, x, y, b):
|
||||
x, y, b = self.virtual_dataset(x, y, b)
|
||||
out = self.gelu(self.matmul1(x, y))
|
||||
out = self.matmul2(out, b)
|
||||
return out
|
||||
|
||||
class Net2(nn.Cell):
|
||||
def __init__(self, strategy1, strategy2, strategy3):
|
||||
super().__init__()
|
||||
self.virtual_dataset = _VirtualDataset()
|
||||
self.get_next = P.GetNext([ms.float32, ms.float32, ms.float32], [[128, 32], [32, 64], [64]], 3, "")
|
||||
self.matmul1 = P.MatMul().shard(strategy1)
|
||||
self.biasadd = P.BiasAdd().shard(strategy2)
|
||||
self.gelu = P.GeLU().shard(strategy3)
|
||||
|
||||
def construct(self, a, b, c):
|
||||
x, y, b = self.get_next()
|
||||
x, y, b = self.virtual_dataset(x, y, b)
|
||||
out = self.gelu(self.matmul1(x, y))
|
||||
out = self.biasadd(out, b)
|
||||
return out
|
||||
|
||||
def compile_net(net, x, y, b):
|
||||
net.set_auto_parallel()
|
||||
net.set_train()
|
||||
_executor.compile(net, x, y, b)
|
||||
|
||||
|
||||
def test_virtual_dataset_model_parallel_semi_auto_parallel():
|
||||
context.set_auto_parallel_context(parallel_mode="semi_auto_parallel")
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0)
|
||||
strategy0 = ((1, 8), (1, 8), (1, 8))
|
||||
context.set_auto_parallel_context(dataset_strategy=strategy0)
|
||||
strategy1 = ((2, 2), (2, 2))
|
||||
strategy2 = ((2, 2), (2, 2))
|
||||
strategy3 = ((2, 4),)
|
||||
net = GradWrap(NetWithLoss(Net1(strategy1, strategy2, strategy3)))
|
||||
x = Tensor(np.ones([128, 32]), dtype=ms.float32)
|
||||
y = Tensor(np.ones([32, 64]), dtype=ms.float32)
|
||||
b = Tensor(np.ones([64, 2048]), dtype=ms.float32)
|
||||
compile_net(net, x, y, b)
|
||||
|
||||
def test_virtual_dataset_model_parallel_auto_parallel():
|
||||
context.set_auto_parallel_context(parallel_mode="auto_parallel")
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0)
|
||||
strategy0 = ((1, 8), (1, 8), (1, 8))
|
||||
context.set_auto_parallel_context(dataset_strategy=strategy0)
|
||||
strategy1 = None
|
||||
strategy2 = None
|
||||
strategy3 = None
|
||||
net = GradWrap(NetWithLoss(Net1(strategy1, strategy2, strategy3)))
|
||||
x = Tensor(np.ones([128, 32]), dtype=ms.float32)
|
||||
y = Tensor(np.ones([32, 64]), dtype=ms.float32)
|
||||
b = Tensor(np.ones([64, 2048]), dtype=ms.float32)
|
||||
compile_net(net, x, y, b)
|
||||
|
||||
def test_virtual_dataset_model_parallel_semi_auto_parallel_diff_input_dim():
|
||||
context.set_auto_parallel_context(parallel_mode="semi_auto_parallel")
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0)
|
||||
strategy0 = ((1, 8), (1, 8), (8,))
|
||||
context.set_auto_parallel_context(dataset_strategy=strategy0)
|
||||
strategy1 = ((2, 2), (2, 2))
|
||||
strategy2 = ((1, 8), (8,))
|
||||
strategy3 = ((2, 4),)
|
||||
x = Tensor(np.ones([128, 32]), dtype=ms.float32)
|
||||
y = Tensor(np.ones([32, 64]), dtype=ms.float32)
|
||||
b = Tensor(np.ones([64]), dtype=ms.float32)
|
||||
net = GradWrap(NetWithLoss(Net2(strategy1, strategy2, strategy3)))
|
||||
compile_net(net, x, y, b)
|
||||
|
||||
def test_virtual_dataset_model_parallel_auto_parallel_diff_input_dim():
|
||||
context.set_auto_parallel_context(parallel_mode="auto_parallel")
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0)
|
||||
strategy0 = ((1, 8), (1, 8), (8,))
|
||||
context.set_auto_parallel_context(dataset_strategy=strategy0)
|
||||
strategy1 = None
|
||||
strategy2 = None
|
||||
strategy3 = None
|
||||
x = Tensor(np.ones([128, 32]), dtype=ms.float32)
|
||||
y = Tensor(np.ones([32, 64]), dtype=ms.float32)
|
||||
b = Tensor(np.ones([64]), dtype=ms.float32)
|
||||
net = GradWrap(NetWithLoss(Net2(strategy1, strategy2, strategy3)))
|
||||
compile_net(net, x, y, b)
|
||||
|
||||
def test_virtual_dataset_model_parallel_semi_auto_parallel_diff_input_dim_not_fully_shard():
|
||||
context.set_auto_parallel_context(parallel_mode="semi_auto_parallel")
|
||||
context.set_auto_parallel_context(device_num=16, global_rank=0)
|
||||
strategy0 = ((1, 8), (1, 8), (1,))
|
||||
context.set_auto_parallel_context(dataset_strategy=strategy0)
|
||||
strategy1 = ((2, 2), (2, 2))
|
||||
strategy2 = ((1, 8), (8,))
|
||||
strategy3 = ((2, 4),)
|
||||
x = Tensor(np.ones([128, 32]), dtype=ms.float32)
|
||||
y = Tensor(np.ones([32, 64]), dtype=ms.float32)
|
||||
b = Tensor(np.ones([64]), dtype=ms.float32)
|
||||
net = GradWrap(NetWithLoss(Net2(strategy1, strategy2, strategy3)))
|
||||
compile_net(net, x, y, b)
|
||||
|
||||
def test_virtual_dataset_model_parallel_auto_parallel_diff_input_dim_not_fully_shard():
|
||||
context.set_auto_parallel_context(parallel_mode="auto_parallel")
|
||||
context.set_auto_parallel_context(device_num=16, global_rank=0)
|
||||
strategy0 = ((1, 8), (1, 8), (1,))
|
||||
context.set_auto_parallel_context(dataset_strategy=strategy0)
|
||||
strategy1 = None
|
||||
strategy2 = None
|
||||
strategy3 = None
|
||||
x = Tensor(np.ones([128, 32]), dtype=ms.float32)
|
||||
y = Tensor(np.ones([32, 64]), dtype=ms.float32)
|
||||
b = Tensor(np.ones([64]), dtype=ms.float32)
|
||||
net = GradWrap(NetWithLoss(Net2(strategy1, strategy2, strategy3)))
|
||||
compile_net(net, x, y, b)
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
context.reset_auto_parallel_context()
|
||||
|
|
@ -131,6 +131,7 @@ def compile_graph_two_input(x, y, net):
|
|||
|
||||
|
||||
def test_dense_relu_semi_auto():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="semi_auto_parallel", full_batch=False)
|
||||
net = DenseMutMulNet()
|
||||
x = Tensor(np.ones([32, 128]).astype(np.float32) * 0.01)
|
||||
|
|
@ -140,6 +141,7 @@ def test_dense_relu_semi_auto():
|
|||
assert v[0][0] == 8
|
||||
|
||||
def test_dense_relu_semi_auto_full_batch():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="semi_auto_parallel", full_batch=True)
|
||||
net = DenseMutMulNet()
|
||||
x = Tensor(np.ones([32, 128]).astype(np.float32) * 0.01)
|
||||
|
|
@ -149,6 +151,7 @@ def test_dense_relu_semi_auto_full_batch():
|
|||
assert v[0][0] == 1
|
||||
|
||||
def test_dense_relu_auto():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="auto_parallel", full_batch=False)
|
||||
net = DenseMutMulNet()
|
||||
x = Tensor(np.ones([32, 128]).astype(np.float32) * 0.01)
|
||||
|
|
@ -158,6 +161,7 @@ def test_dense_relu_auto():
|
|||
assert v[0][0] == 8
|
||||
|
||||
def test_dense_relu_auto_full_batch():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="auto_parallel", full_batch=True)
|
||||
net = DenseMutMulNet()
|
||||
x = Tensor(np.ones([32, 128]).astype(np.float32) * 0.01)
|
||||
|
|
@ -167,6 +171,7 @@ def test_dense_relu_auto_full_batch():
|
|||
assert v[0][0] == 1
|
||||
|
||||
def test_mul_neg_two_output_semi_auto():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="semi_auto_parallel", full_batch=False)
|
||||
net = MulNegTwoOutputNet()
|
||||
x = Tensor(np.ones([32, 128]).astype(np.float32) * 0.01)
|
||||
|
|
@ -179,6 +184,7 @@ def test_mul_neg_two_output_semi_auto():
|
|||
assert count == 2
|
||||
|
||||
def test_mul_neg_two_output_semi_auto_full_batch():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="semi_auto_parallel", full_batch=True)
|
||||
net = MulNegTwoOutputNet()
|
||||
x = Tensor(np.ones([32, 128]).astype(np.float32) * 0.01)
|
||||
|
|
@ -191,6 +197,7 @@ def test_mul_neg_two_output_semi_auto_full_batch():
|
|||
assert count == 2
|
||||
|
||||
def test_mul_neg_two_output_auto():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="auto_parallel", full_batch=False)
|
||||
net = MulNegTwoOutputNet()
|
||||
x = Tensor(np.ones([32, 128]).astype(np.float32) * 0.01)
|
||||
|
|
@ -203,6 +210,7 @@ def test_mul_neg_two_output_auto():
|
|||
assert count == 2
|
||||
|
||||
def test_mul_neg_two_output_full_batch():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="auto_parallel", full_batch=True)
|
||||
net = MulNegTwoOutputNet()
|
||||
x = Tensor(np.ones([32, 128]).astype(np.float32) * 0.01)
|
||||
|
|
@ -215,6 +223,7 @@ def test_mul_neg_two_output_full_batch():
|
|||
assert count == 2
|
||||
|
||||
def test_reshape_matmul_semi_auto():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="semi_auto_parallel", full_batch=False)
|
||||
strategy1 = None
|
||||
strategy2 = ((1, 1), (1, 8))
|
||||
|
|
@ -226,6 +235,7 @@ def test_reshape_matmul_semi_auto():
|
|||
assert v[0][0] == 8
|
||||
|
||||
def test_reshape_matmul_auto():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="auto_parallel", full_batch=False)
|
||||
strategy1 = None
|
||||
strategy2 = ((1, 1), (1, 8))
|
||||
|
|
@ -237,6 +247,7 @@ def test_reshape_matmul_auto():
|
|||
assert v[0][0] == 8
|
||||
|
||||
def test_matmul_reshape_semi_auto():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="semi_auto_parallel", full_batch=False)
|
||||
strategy2 = None
|
||||
strategy1 = ((1, 1), (1, 8))
|
||||
|
|
@ -248,6 +259,7 @@ def test_matmul_reshape_semi_auto():
|
|||
assert v[0][0] == 8
|
||||
|
||||
def test_matmul_reshape_auto():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="auto_parallel", full_batch=False)
|
||||
strategy2 = None
|
||||
strategy1 = ((1, 1), (1, 8))
|
||||
|
|
@ -259,6 +271,7 @@ def test_matmul_reshape_auto():
|
|||
assert v[0][0] == 8
|
||||
|
||||
def test_reshape_mul_semi_auto():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="semi_auto_parallel", full_batch=True)
|
||||
net = ReshapeMulNet()
|
||||
x = Tensor(np.ones([64, 4]), ms.float32)
|
||||
|
|
@ -268,6 +281,7 @@ def test_reshape_mul_semi_auto():
|
|||
assert v[0][0] == 1
|
||||
|
||||
def test_reshape_mul_auto():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="auto_parallel", full_batch=True)
|
||||
net = ReshapeMulNet()
|
||||
x = Tensor(np.ones([64, 4]), ms.float32)
|
||||
|
|
@ -277,6 +291,7 @@ def test_reshape_mul_auto():
|
|||
assert v[0][0] == 1
|
||||
|
||||
def test_scalar_output_semi_auto():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="semi_auto_parallel", full_batch=False)
|
||||
net = ParallelMulNet()
|
||||
loss_fn = nn.SoftmaxCrossEntropyWithLogits(reduction='mean')
|
||||
|
|
@ -292,6 +307,7 @@ def test_scalar_output_semi_auto():
|
|||
assert count == 1
|
||||
|
||||
def test_scalar_output_auto():
|
||||
context.reset_auto_parallel_context()
|
||||
context.set_auto_parallel_context(device_num=8, global_rank=0, parallel_mode="auto_parallel", full_batch=False)
|
||||
net = ParallelMulNet()
|
||||
loss_fn = nn.SoftmaxCrossEntropyWithLogits(reduction='mean')
|
||||
|
|
|
|||
Loading…
Reference in New Issue