| 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. |
nothing calls this directly
no test coverage detected