mindspore代码评注-等风也等你 #28

Open
darrenlu wants to merge 8 commits from darrenlu/mindspore2022:master into master
1 changed files with 23 additions and 7 deletions
Showing only changes of commit 88f3bf423e - Show all commits

View File

@ -37,29 +37,38 @@ inline const PrimitivePtr kPrimGkDropout = std::make_shared<Primitive>("GkDropou
namespace graphkernel {
using opt::CheckCNodeInputSize;
using opt::kDropoutInputTensorNum;
// 静态成员变量,用于保存随机数生成器的种子,初始化为当前时间戳
int64_t DropoutExpander::seed_ = time(nullptr);
// 预处理函数用于处理Dropout节点之前的操作
AnfNodePtr DropoutExpander::PreProcess(const FuncGraphPtr &func_graph, const AnfNodePtr &node) {
// 检查节点和节点类型是否有效
MS_EXCEPTION_IF_NULL(node);
CNodePtr cnode = node->cast<CNodePtr>();
MS_EXCEPTION_IF_NULL(cnode);
// 检查Dropout节点输入的数量是否正确
CheckCNodeInputSize(cnode, kDropoutInputTensorNum);
// 获取输入节点的设备上的形状
auto shape = AnfAlgo::GetInputDeviceShape(cnode, 0);
ShapeVector shape_i64;
(void)std::transform(shape.begin(), shape.end(), std::back_inserter(shape_i64), SizeToLong);
// Get seed from original dropout's attrs, rather than set seed by time.
// Only seed0 and seed1 are all equal to 0, then set seed = time.
// 从Dropout节点的属性中获取种子值
auto node_prim = GetCNodePrimitive(node);
MS_EXCEPTION_IF_NULL(node_prim);
int64_t seed = GetValue<int64_t>(node_prim->GetAttr("Seed0"));
// 如果Seed0和Seed1都等于0则使用全局种子值seed_
if (seed == 0) {
seed = GetValue<int64_t>(node_prim->GetAttr("Seed1"));
if (seed == 0) {
seed = seed_++;
}
}
// Create a uniform_real kernel to generate random value.
// 创建一个uniform_real节点用于生成随机值
auto tensor = std::make_shared<tensor::Tensor>(kNumberTypeInt64, ShapeVector(1, SizeToLong(shape.size())),
static_cast<void *>(&shape[0]), kNumberTypeInt64);
AnfNodePtrList uniform_real_input = {NewValueNode(prim::kPrimCudnnUniformReal), NewValueNode(tensor)};
@ -69,7 +78,8 @@ AnfNodePtr DropoutExpander::PreProcess(const FuncGraphPtr &func_graph, const Anf
SetNodeAttrSafely("seed", MakeValue(seed), uniform_real_node);
common::AnfAlgo::SetNodeAttr("seed2", MakeValue(static_cast<int64_t>(0)), uniform_real_node);
uniform_real_node->set_abstract(std::make_shared<abstract::AbstractTensor>(kFloat32, shape_i64));
// Set kernel_info for uniform_real node
// 为uniform_real节点设置kernel_info
auto uniform_real_kernel_info_builder = std::make_shared<kernel::KernelBuildInfo::KernelBuildInfoBuilder>();
uniform_real_kernel_info_builder->SetInputsFormat({kOpFormat_DEFAULT});
uniform_real_kernel_info_builder->SetInputsDeviceType({kNumberTypeInt32});
@ -79,23 +89,29 @@ AnfNodePtr DropoutExpander::PreProcess(const FuncGraphPtr &func_graph, const Anf
uniform_real_kernel_info_builder->SetProcessor(kernel::Processor::CUDA);
AnfAlgo::SetSelectKernelBuildInfo(uniform_real_kernel_info_builder->Build(), uniform_real_node.get());
// Create a GKDropout node with uniform_real as its second input.
// 创建一个GKDropout节点其第二个输入为uniform_real节点
AnfNodePtrList gkdropout_inputs = {NewValueNode(prim::kPrimGkDropout), cnode->input(1), uniform_real_node};
auto new_dropout_node = func_graph->NewCNode(gkdropout_inputs);
SetNodeAttrSafely("keep_prob", MakeValue(common::AnfAlgo::GetNodeAttr<float>(cnode, "keep_prob")), new_dropout_node);
// the output info is unchanged.
// 输出信息与原始节点相同
new_dropout_node->set_abstract(node->abstract());
auto old_kernel_info = AnfAlgo::GetSelectKernelBuildInfo(node);
auto dropout_kernel_info_builder = std::make_shared<kernel::KernelBuildInfo::KernelBuildInfoBuilder>(old_kernel_info);
dropout_kernel_info_builder->SetInputsFormat({old_kernel_info->GetInputFormat(0), kOpFormat_DEFAULT});
dropout_kernel_info_builder->SetInputsDeviceType({old_kernel_info->GetInputDeviceType(0), kNumberTypeFloat32});
AnfAlgo::SetSelectKernelBuildInfo(dropout_kernel_info_builder->Build(), new_dropout_node.get());
return new_dropout_node;
}
// 运行函数用于处理Dropout节点
AnfNodePtr DropoutExpander::Run(const AnfNodePtr &node) {
// 预处理Dropout节点
auto gkdropout_node = PreProcess(node->func_graph(), node);
// 调用父类的Run方法处理GKDropout节点
return PyExpander::Run(gkdropout_node);
}
} // namespace graphkernel
} // namespace mindspore