MCPcopy Create free account
hub / github.com/conflow-dev/ConFlow / Compute

Method Compute

en_ops/e_huber_loss.cc:99–132  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

97 OP_REQUIRES(context, lib != NULL, errors::Unknown("Unable to load sgx.so!"));
98 }
99 void Compute(OpKernelContext *context) override
100 {
101
102 const Tensor &grad = context->input(0);
103 const Tensor &input = context->input(1);
104 const Tensor &real = context->input(2);
105 const Tensor &delta = context->input(3);
106 auto input_flat = input.flat<float>();
107 auto real_flat = real.flat<float>();
108 auto grad_flat = grad.flat<float>();
109 auto delta_flat = delta.flat<float>();
110
111 const TensorShape &input_shape = input.shape();
112
113 Tensor *output = NULL;
114 OP_REQUIRES_OK(context, context->allocate_output(0, input_shape, &output));
115 auto output_flat = output->flat<float>();
116
117 Tensor *grad_r = NULL;
118 OP_REQUIRES_OK(context, context->allocate_output(1, real.shape(), &grad_r));
119 auto grad_r_flat = grad_r->flat<float>();
120
121 int N = input_flat.size() / times_;
122 int C = delta_flat.size();
123 int B = input_shape.dim_size(0) / times_;
124
125 unsigned long int eid_ = (eid_high_ << 32) + eid_low_;
126 typedef void (*function)(unsigned long int eid, float *pred, float *real, float *delta, float *grad, int N, int C, int B, float *output, float *grad_r);
127 dlerror();
128 function huberloss_grad_kernel = (function)dlsym(lib, "huberloss_grad");
129 const char *dlsym_error = dlerror();
130 OP_REQUIRES(context, !dlsym_error, errors::Unknown("loading of huberloss_grad failed: ", dlsym_error));
131 huberloss_grad_kernel(eid_, (float *)input_flat.data(), (float *)real_flat.data(), (float *)delta_flat.data(), (float *)grad_flat.data(), N, C, B, (float *)output_flat.data(), (float *)grad_r_flat.data());
132 };
133
134private:
135 void *lib;

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected