Add some nullptr checks (#160)

This commit is contained in:
Changming Sun 2018-12-12 16:53:35 -08:00 committed by GitHub
parent 618cc51754
commit ed98d3d653
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
5 changed files with 10 additions and 7 deletions

View file

@ -48,7 +48,9 @@ static Status ComputeRange(OpKernelContext* ctx) {
}
Status Range::Compute(OpKernelContext* ctx) const {
auto data_type = ctx->Input<Tensor>(0)->DataType();
auto input_tensor = ctx->Input<Tensor>(0);
if (input_tensor == nullptr) return Status(common::ONNXRUNTIME, common::FAIL, "input count mismatch");
auto data_type = input_tensor->DataType();
if (data_type == DataTypeImpl::GetType<int32_t>()) {
return ComputeRange<int32_t>(ctx);
}

View file

@ -207,6 +207,7 @@ Status StringNormalizer::Compute(OpKernelContext* ctx) const {
using namespace string_normalizer;
auto X = ctx->Input<Tensor>(0);
if (X == nullptr) return Status(common::ONNXRUNTIME, common::FAIL, "input count mismatch");
auto& input_dims = X->Shape().GetDims();
size_t N = 0;

View file

@ -450,6 +450,7 @@ Status Tokenizer::SeparatorTokenize(OpKernelContext* ctx,
Status Tokenizer::Compute(OpKernelContext* ctx) const {
// Get input buffer ptr
auto X = ctx->Input<Tensor>(0);
if (X == nullptr) return Status(common::ONNXRUNTIME, common::FAIL, "input count mismatch");
if (X->DataType() != DataTypeImpl::GetType<std::string>()) {
return Status(common::ONNXRUNTIME, common::INVALID_ARGUMENT,
"tensor(string) expected as input");

View file

@ -33,19 +33,18 @@ Status EyeLike::Compute(OpKernelContext* context) const {
auto output_tensor_dtype = has_dtype_ ? static_cast<onnx::TensorProto::DataType>(dtype_) : utils::GetTensorProtoType(*T1);
switch (output_tensor_dtype) {
case onnx::TensorProto_DataType_FLOAT:
return ComputeImpl<float>(context);
return ComputeImpl<float>(context, T1);
case onnx::TensorProto_DataType_INT64:
return ComputeImpl<int64_t>(context);
return ComputeImpl<int64_t>(context, T1);
case onnx::TensorProto_DataType_UINT64:
return ComputeImpl<uint64_t>(context);
return ComputeImpl<uint64_t>(context, T1);
default:
ONNXRUNTIME_THROW("Unsupported 'dtype' value: ", output_tensor_dtype);
}
}
template <typename T>
Status EyeLike::ComputeImpl(OpKernelContext* context) const {
const Tensor* T1 = context->Input<Tensor>(0);
Status EyeLike::ComputeImpl(OpKernelContext* context, const Tensor* T1) const {
const std::vector<int64_t>& input_dims = T1->Shape().GetDims();
if (input_dims.size() != 2) {
return Status(ONNXRUNTIME, INVALID_ARGUMENT, "EyeLike : Input tensor dimension is not 2");

View file

@ -22,7 +22,7 @@ class EyeLike final : public OpKernel {
private:
template <typename T>
Status ComputeImpl(OpKernelContext* context) const;
Status ComputeImpl(OpKernelContext* context, const Tensor* T1) const;
bool has_dtype_;
int64_t dtype_;