MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / Compute

Method Compute

tensorflow/core/kernels/mkl_lrn_op.cc:340–456  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

338 }
339
340 void Compute(OpKernelContext* context) override {
341 try {
342 SanityCheckInputs(context);
343 if (!context->status().ok()) return;
344
345 MklDnnData<T> input_grad_dnn_data(&cpu_engine_);
346 MklDnnData<T> orig_input_dnn_data(&cpu_engine_);
347 MklDnnData<T> orig_output_dnn_data(&cpu_engine_);
348 MklDnnData<T> output_dnn_data(&cpu_engine_);
349
350 MklDnnShape input_grad_dnn_shape, orig_input_dnn_shape,
351 orig_output_dnn_shape;
352 GetMklShape(context, kIdxGradient, &input_grad_dnn_shape);
353 GetMklShape(context, kIdxOrigInput, &orig_input_dnn_shape);
354 GetMklShape(context, kIdxOrigOutput, &orig_output_dnn_shape);
355
356 // We only use DNNL if all of the necessary inputs are present
357 // in dnnl format, and Channel is the last dimension
358 bool can_use_dnnl = workspace_enabled_ &&
359 input_grad_dnn_shape.IsMklTensor() &&
360 orig_input_dnn_shape.IsMklTensor() &&
361 orig_output_dnn_shape.IsMklTensor() &&
362 input_grad_dnn_shape.IsMklChannelDim(
363 input_grad_dnn_shape.GetDimension() - 1) &&
364 orig_input_dnn_shape.IsMklChannelDim(
365 orig_input_dnn_shape.GetDimension() - 1) &&
366 orig_output_dnn_shape.IsMklChannelDim(
367 orig_output_dnn_shape.GetDimension() - 1);
368
369 if (!can_use_dnnl) {
370 // Fallback to eigen
371 MklDefaultToEigen(context);
372 return;
373 }
374 // At this point, we have the all clear to use OneDNN constructs
375 // Naming: diff_dst is input_gradient_tensor; src is orig_input_tensor.
376 const Tensor& input_grad_tensor = MklGetInput(context, kIdxGradient);
377 const Tensor& orig_input_tensor = MklGetInput(context, kIdxOrigInput);
378
379 // Get input sizes in OneDNN required NCHW format.
380 // LRN does not have data_format attribute. But by default it has
381 // NHWC format.
382 memory::desc original_output_md = orig_output_dnn_shape.GetCurLayout();
383 memory::desc target_diff_dst_md = ConfigureInputGradient(
384 input_grad_tensor, input_grad_dnn_shape, &input_grad_dnn_data);
385
386 memory::desc orig_input_md = orig_input_dnn_shape.GetCurLayout();
387 memory::dims orig_input_dims =
388 orig_input_dnn_shape.GetSizesAsMklDnnDims();
389 orig_input_dnn_data.SetUsrMem(orig_input_md, &orig_input_tensor);
390 orig_input_dnn_data.SetOpMemDesc(orig_input_dims, MEMORY_FORMAT::nhwc);
391 orig_input_dnn_data.SetUsrMemDataHandle(&orig_input_tensor, bwd_stream_);
392
393 // output_dnn_data has the same shape as original input
394 output_dnn_data.SetUsrMem(orig_input_md);
395 output_dnn_data.SetOpMemDesc(orig_input_dims, MEMORY_FORMAT::nhwc);
396
397 // OneDNN has a notion of kernel_size and not depth_radius.

Callers

nothing calls this directly

Calls 15

GetMklShapeFunction · 0.85
to_stringFunction · 0.85
IsMklTensorMethod · 0.80
IsMklChannelDimMethod · 0.80
GetCurLayoutMethod · 0.80
GetSizesAsMklDnnDimsMethod · 0.80
SetUsrMemMethod · 0.80
SetOpMemDescMethod · 0.80
SetUsrMemDataHandleMethod · 0.80
GetTfDataFormatMethod · 0.80
CheckReorderToOpMemMethod · 0.80
okMethod · 0.45

Tested by

no test coverage detected