From b1e83b33b001d9f046ba6569362ad1fd4c28ebf6 Mon Sep 17 00:00:00 2001 From: Zeng Jinle <32832641+sneaxiy@users.noreply.github.com> Date: Tue, 24 Sep 2019 10:19:47 +0800 Subject: [PATCH] fix huber loss op attr type, test=develop (#19937) --- paddle/fluid/operators/huber_loss_op.h | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/paddle/fluid/operators/huber_loss_op.h b/paddle/fluid/operators/huber_loss_op.h index fa21bd01cb0..7000b5d3acc 100644 --- a/paddle/fluid/operators/huber_loss_op.h +++ b/paddle/fluid/operators/huber_loss_op.h @@ -41,7 +41,7 @@ struct HuberLossForward { T delta; }; -template +template class HuberLossKernel : public framework::OpKernel { public: void Compute(const framework::ExecutionContext& context) const override { @@ -49,7 +49,7 @@ class HuberLossKernel : public framework::OpKernel { auto* in1 = context.Input("Y"); auto* out0 = context.Output("Residual"); auto* out1 = context.Output("Out"); - auto delta = static_cast(context.Attr("delta")); + auto delta = static_cast(context.Attr("delta")); auto& place = *context.template device_context().eigen_device(); @@ -86,7 +86,7 @@ struct HuberLossBackward { T delta; }; -template +template class HuberLossGradKernel : public framework::OpKernel { public: void Compute(const framework::ExecutionContext& context) const override { @@ -94,7 +94,7 @@ class HuberLossGradKernel : public framework::OpKernel { auto* in1 = context.Input(framework::GradVarName("Out")); auto* out0 = context.Output(framework::GradVarName("X")); auto* out1 = context.Output(framework::GradVarName("Y")); - auto delta = static_cast(context.op().Attr("delta")); + auto delta = static_cast(context.op().Attr("delta")); auto& place = *context.template device_context().eigen_device(); -- GitLab