From 55991822a09ac1d675899567ff867cd6feff79c4 Mon Sep 17 00:00:00 2001 From: xzl Date: Fri, 8 Sep 2017 22:51:14 +0800 Subject: [PATCH] modify GetAttr to Attr --- paddle/operators/transpose_op.cc | 2 +- paddle/operators/transpose_op.cu | 4 ++-- paddle/operators/transpose_op.h | 4 ++-- 3 files changed, 5 insertions(+), 5 deletions(-) diff --git a/paddle/operators/transpose_op.cc b/paddle/operators/transpose_op.cc index 9b7812c79..ea6b2a9ec 100644 --- a/paddle/operators/transpose_op.cc +++ b/paddle/operators/transpose_op.cc @@ -28,7 +28,7 @@ class TransposeOp : public framework::OperatorWithKernel { protected: void InferShape(const framework::InferShapeContext &ctx) const override { auto in_dim = ctx.Input("X")->dims(); - auto axis = ctx.GetAttr>("axis"); + auto axis = ctx.Attr>("axis"); size_t in_dim_size = in_dim.size(); size_t axis_size = axis.size(); diff --git a/paddle/operators/transpose_op.cu b/paddle/operators/transpose_op.cu index 853659e3c..24feeea4b 100644 --- a/paddle/operators/transpose_op.cu +++ b/paddle/operators/transpose_op.cu @@ -98,7 +98,7 @@ class TransposeCUDAKernel : public framework::OpKernel { "It must use GPUPlace."); auto* in = context.Input("X"); auto* out = context.Output("Out"); - auto axis = context.GetAttr>("axis"); + auto axis = context.Attr>("axis"); TransposeCUDA(context, *in, *out, axis); } }; @@ -111,7 +111,7 @@ class TransposeGradCUDAKernel : public framework::OpKernel { "It must use GPUPlace."); auto* in = context.Input(framework::GradVarName("Out")); auto* out = context.Output(framework::GradVarName("X")); - auto axis_temp = context.GetAttr>("axis"); + auto axis_temp = context.Attr>("axis"); std::vector axis(axis_temp); diff --git a/paddle/operators/transpose_op.h b/paddle/operators/transpose_op.h index ca64b5a63..57f63e60e 100644 --- a/paddle/operators/transpose_op.h +++ b/paddle/operators/transpose_op.h @@ -77,7 +77,7 @@ class TransposeKernel : public framework::OpKernel { auto* out = context.Output("Out"); out->mutable_data(context.GetPlace()); - auto axis = context.GetAttr>("axis"); + auto axis = context.Attr>("axis"); int ndims = axis.size(); switch (ndims) { case 2: @@ -107,7 +107,7 @@ class TransposeGradKernel : public framework::OpKernel { auto* out = context.Output(framework::GradVarName("X")); out->mutable_data(context.GetPlace()); - auto axis_temp = context.GetAttr>("axis"); + auto axis_temp = context.Attr>("axis"); std::vector axis(axis_temp); for (size_t i = 0; i < axis.size(); i++) { -- GitLab