mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
DML functions always returning a value (#9485)
* Always return a value * @fdwr advice added
This commit is contained in:
parent
a2b3e6bb23
commit
2d44bd525b
11 changed files with 37 additions and 10 deletions
|
|
@ -28,6 +28,7 @@ onnx::OpSchema::FormalParameterOption AbiCustomRegistry::ConvertFormalParameterO
|
|||
|
||||
default:
|
||||
THROW_HR(E_NOTIMPL);
|
||||
return onnx::OpSchema::FormalParameterOption::Single;
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -87,7 +87,9 @@ MLOperatorTensorDataType GetMlDataTypeFromDmlDataType(DML_TENSOR_DATA_TYPE tenso
|
|||
case DML_TENSOR_DATA_TYPE_INT64: return MLOperatorTensorDataType::Int64;
|
||||
case DML_TENSOR_DATA_TYPE_FLOAT64: return MLOperatorTensorDataType::Double;
|
||||
|
||||
default: ML_INVALID_ARGUMENT("Unknown DML_TENSOR_DATA_TYPE.");
|
||||
default:
|
||||
ML_INVALID_ARGUMENT("Unknown DML_TENSOR_DATA_TYPE.");
|
||||
return MLOperatorTensorDataType::Undefined;
|
||||
};
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -56,7 +56,9 @@ struct ActivationOperatorDesc
|
|||
case DML_OPERATOR_ACTIVATION_TANH: return { activationType, ¶ms.tanh };
|
||||
case DML_OPERATOR_ACTIVATION_THRESHOLDED_RELU: return { activationType, ¶ms.thresholdedRelu };
|
||||
case DML_OPERATOR_ACTIVATION_SHRINK: return { activationType, ¶ms.shrink };
|
||||
default: THROW_HR(E_INVALIDARG);
|
||||
default:
|
||||
THROW_HR(E_INVALIDARG);
|
||||
return { activationType, ¶ms.relu };
|
||||
}
|
||||
}
|
||||
};
|
||||
|
|
@ -206,9 +208,9 @@ private:
|
|||
|
||||
~DynamicBucket()
|
||||
{
|
||||
if (data)
|
||||
if (this->data)
|
||||
{
|
||||
(void)VirtualFree(data, 0, MEM_RELEASE);
|
||||
(void)VirtualFree(this->data, 0, MEM_RELEASE);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
|
|
|||
|
|
@ -2159,6 +2159,7 @@ auto OperatorTypeVisitor(DML_OPERATOR_TYPE type, Visitor&& visitor, Ts&&... args
|
|||
|
||||
default:
|
||||
THROW_HR(E_INVALIDARG);
|
||||
return std::invoke(std::forward<Visitor>(visitor), DML_ACTIVATION_RELU_OPERATOR_DESC{}, std::forward<Ts>(args)...);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1484,7 +1484,9 @@ inline const DML_OPERATOR_SCHEMA& GetSchema(DML_OPERATOR_TYPE operatorType)
|
|||
case DML_OPERATOR_ACTIVATION_THRESHOLDED_RELU: return DML_ACTIVATION_THRESHOLDED_RELU_OPERATOR_SCHEMA;
|
||||
case DML_OPERATOR_ACTIVATION_SHRINK: return DML_ACTIVATION_SHRINK_OPERATOR_SCHEMA;
|
||||
|
||||
default: THROW_HR(E_INVALIDARG);
|
||||
default:
|
||||
THROW_HR(E_INVALIDARG);
|
||||
return DML_ACTIVATION_RELU_OPERATOR_SCHEMA;
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -2052,7 +2054,11 @@ inline AbstractOperatorDesc ConvertOperatorDesc(const DML_OPERATOR_DESC& opDesc)
|
|||
return AbstractOperatorDesc(
|
||||
&DML_ACTIVATION_SHRINK_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_ACTIVATION_SHRINK_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
default: THROW_HR(E_INVALIDARG);
|
||||
default:
|
||||
THROW_HR(E_INVALIDARG);
|
||||
return AbstractOperatorDesc(
|
||||
&DML_ACTIVATION_RELU_OPERATOR_SCHEMA,
|
||||
GetFields(*static_cast<const DML_ACTIVATION_RELU_OPERATOR_DESC*>(opDesc.Desc)));
|
||||
}
|
||||
|
||||
}
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ namespace Dml::GraphDescBuilder
|
|||
|
||||
assert(false);
|
||||
THROW_HR(E_UNEXPECTED);
|
||||
return node.OutputDefs()[0]->Name();
|
||||
}
|
||||
|
||||
GraphDesc BuildGraphDesc(
|
||||
|
|
|
|||
|
|
@ -43,6 +43,7 @@ size_t AttributeValue::ElementCount() const {
|
|||
// The type is validated when default attributes are registered
|
||||
assert(false);
|
||||
THROW_HR(E_FAIL);
|
||||
return 0;
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -238,6 +239,7 @@ struct MLTypeTraits<onnxruntime::MLFloat16> {
|
|||
ML_TENSOR_TYPE_CASE(onnxruntime::MLFloat16);
|
||||
|
||||
THROW_HR(E_NOTIMPL);
|
||||
return MLOperatorTensorDataType::Undefined;
|
||||
}
|
||||
|
||||
#undef ML_TENSOR_TYPE_CASE
|
||||
|
|
@ -264,6 +266,7 @@ onnxruntime::MLDataType ToTensorDataType(::MLOperatorTensorDataType type) {
|
|||
ML_TENSOR_TYPE_CASE(onnxruntime::MLFloat16);
|
||||
|
||||
THROW_HR(E_NOTIMPL);
|
||||
return onnxruntime::DataTypeImpl::GetTensorType<float>();
|
||||
}
|
||||
|
||||
::MLOperatorTensorDataType ToMLTensorDataType(onnx::TensorProto_DataType type) {
|
||||
|
|
@ -315,6 +318,7 @@ onnxruntime::MLDataType ToTensorDataType(::MLOperatorTensorDataType type) {
|
|||
|
||||
default:
|
||||
THROW_HR(E_NOTIMPL);
|
||||
return MLOperatorTensorDataType::Undefined;
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -389,6 +393,7 @@ std::string ToTypeString(MLOperatorEdgeDescription desc) {
|
|||
|
||||
default:
|
||||
THROW_HR(E_NOTIMPL);
|
||||
return "";
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -594,7 +599,7 @@ HRESULT OpNodeInfoWrapper<NodeInfoImpl_t, Base1_t, Base2_t>::GetAttributeHelper(
|
|||
uint32_t elementByteSize,
|
||||
void* value) const {
|
||||
using elementType_t = typename MLAttributeTypeTraits<T>::Type;
|
||||
static_assert(!typename MLAttributeTypeTraits<T>::IsArray, "This function only works for simple non-array types.");
|
||||
static_assert(!MLAttributeTypeTraits<T>::IsArray, "This function only works for simple non-array types.");
|
||||
ML_CHECK_BOOL(sizeof(elementType_t) == elementByteSize);
|
||||
THROW_IF_NOT_OK(m_impl->template GetAttr<elementType_t>(name, static_cast<elementType_t*>(value)));
|
||||
return S_OK;
|
||||
|
|
|
|||
|
|
@ -62,6 +62,7 @@ public:
|
|||
if (direction == AttrValue::DirectionBidirectional) { return DML_RECURRENT_NETWORK_DIRECTION_BIDIRECTIONAL; }
|
||||
|
||||
ML_INVALID_ARGUMENT("Unsupported direction"); // throws
|
||||
return DML_RECURRENT_NETWORK_DIRECTION_FORWARD;
|
||||
}
|
||||
|
||||
void InitActivationDescs(const MLOperatorKernelCreationContext& kernelInfo, _Out_ std::vector<DML_OPERATOR_DESC>& descs, gsl::span<const std::string> defaultActivations)
|
||||
|
|
|
|||
|
|
@ -436,6 +436,7 @@ namespace Dml
|
|||
return *index;
|
||||
}
|
||||
ML_INVALID_ARGUMENT("Unknown interpolation mode");
|
||||
return (DML_INTERPOLATION_MODE)0;
|
||||
}
|
||||
|
||||
DML_DEPTH_SPACE_ORDER MapStringToDepthSpaceMode(std::string_view mode)
|
||||
|
|
@ -450,6 +451,7 @@ namespace Dml
|
|||
return *index;
|
||||
}
|
||||
ML_INVALID_ARGUMENT("Unknown depth/space order");
|
||||
return (DML_DEPTH_SPACE_ORDER)0;
|
||||
}
|
||||
|
||||
} // namespace Dml
|
||||
|
|
|
|||
|
|
@ -119,7 +119,9 @@ inline size_t GetByteSizeFromMlDataType(MLOperatorTensorDataType tensorDataType)
|
|||
case MLOperatorTensorDataType::Complex64: return 8;
|
||||
case MLOperatorTensorDataType::Complex128: return 16;
|
||||
case MLOperatorTensorDataType::Undefined:
|
||||
default: THROW_HR(E_INVALIDARG);
|
||||
default:
|
||||
THROW_HR(E_INVALIDARG);
|
||||
return 0;
|
||||
};
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -163,7 +163,9 @@ namespace OperatorHelper
|
|||
case MLOperatorTensorDataType::Complex64: return static_cast<int64_t>(*reinterpret_cast<const float*>(p)); // Read the real component.
|
||||
case MLOperatorTensorDataType::Complex128: return static_cast<int64_t>(*reinterpret_cast<const double*>(p)); // Read the real component.
|
||||
case MLOperatorTensorDataType::Undefined:
|
||||
default: ML_INVALID_ARGUMENT("Unknown MLOperatorTensorDataType.");
|
||||
default:
|
||||
ML_INVALID_ARGUMENT("Unknown MLOperatorTensorDataType.");
|
||||
return 0;
|
||||
};
|
||||
}
|
||||
|
||||
|
|
@ -187,7 +189,9 @@ namespace OperatorHelper
|
|||
case MLOperatorTensorDataType::Complex64: return static_cast<double>(*reinterpret_cast<const float*>(p)); // Read the real component.
|
||||
case MLOperatorTensorDataType::Complex128: return static_cast<double>(*reinterpret_cast<const double*>(p)); // Read the real component.
|
||||
case MLOperatorTensorDataType::Undefined:
|
||||
default: ML_INVALID_ARGUMENT("Unknown MLOperatorTensorDataType.");
|
||||
default:
|
||||
ML_INVALID_ARGUMENT("Unknown MLOperatorTensorDataType.");
|
||||
return 0.0;
|
||||
};
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue