mindspore代码评注-等风也等你 #28
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Reference in New Issue