mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-20 19:12:24 +00:00
Conv node bug, cached state was incoherent (#10041)
* Moved the init earlier to keep the cache coherent * Move setting of w_desc later, and zero shape check later to catch all cacheable changes. * Add comment
This commit is contained in:
parent
f4b2d3af2b
commit
c1cf16ed5d
1 changed files with 6 additions and 4 deletions
|
|
@ -176,9 +176,6 @@ Status Conv<T>::UpdateState(OpKernelContext* context, bool bias_expected) const
|
|||
s_.slice_axes = slice_axes;
|
||||
|
||||
s_.Y = context->Output(0, TensorShape(s_.y_dims));
|
||||
if (s_.Y->Shape().Size() == 0) {
|
||||
return Status::OK();
|
||||
}
|
||||
if (post_slicing_required) {
|
||||
// Post slicing needed. Create and fill in the Conv results in an intermediate buffer.
|
||||
s_.memory_for_cudnn_conv_results = GetScratchBuffer<void>(TensorShape(y_dims_with_adjusted_pads).Size() * s_.element_size);
|
||||
|
|
@ -206,9 +203,14 @@ Status Conv<T>::UpdateState(OpKernelContext* context, bool bias_expected) const
|
|||
dilations.push_back(1);
|
||||
}
|
||||
|
||||
if (w_dims_changed) {
|
||||
if (w_dims_changed)
|
||||
ORT_RETURN_IF_ERROR(s_.w_desc.Set(w_dims, CudnnTensor::GetDataType<CudaT>()));
|
||||
|
||||
// We must delay returning early until here so that the weight dims have been cached properly
|
||||
if (s_.Y->Shape().Size() == 0) {
|
||||
return Status::OK();
|
||||
}
|
||||
|
||||
ORT_RETURN_IF_ERROR(s_.x_tensor.Set(x_dims_cudnn, CudnnTensor::GetDataType<CudaT>()));
|
||||
ORT_RETURN_IF_ERROR(s_.y_tensor.Set(y_dims_cudnn, CudnnTensor::GetDataType<CudaT>()));
|
||||
ORT_RETURN_IF_ERROR(s_.conv_desc.Set(kernel_shape.size(), pads, strides, dilations,
|
||||
|
|
|
|||
Loading…
Reference in a new issue