forked from mindspore-Ecosystem/mindspore
!34345 TensorScatterUpdate bool bug
Merge pull request !34345 from ling/pr
This commit is contained in:
commit
9ae8279e40
|
|
@ -29,6 +29,9 @@ namespace mindspore::ops {
|
|||
const std::set<TypePtr> common_valid_types = {kInt8, kInt16, kInt32, kInt64, kUInt8, kUInt16,
|
||||
kUInt32, kUInt64, kFloat16, kFloat32, kFloat64};
|
||||
|
||||
const std::set<TypePtr> common_valid_types_with_bool = {kInt8, kInt16, kInt32, kInt64, kUInt8, kUInt16,
|
||||
kUInt32, kUInt64, kFloat16, kFloat32, kFloat64, kBool};
|
||||
|
||||
const std::set<TypePtr> common_valid_types_with_complex = {kInt8, kInt16, kInt32, kInt64, kUInt8,
|
||||
kUInt16, kUInt32, kUInt64, kFloat16, kFloat32,
|
||||
kFloat64, kComplex64, kComplex128};
|
||||
|
|
|
|||
|
|
@ -83,6 +83,9 @@ TypePtr TensorScatterArithmeticInferType(const PrimitivePtr &primitive,
|
|||
std::map<std::string, TypePtr> type_dict;
|
||||
type_dict.emplace("input_x", input_args[kInputIndex0]->BuildType());
|
||||
type_dict.emplace("updates", input_args[kInputIndex2]->BuildType());
|
||||
if (prim_name == prim::kPrimTensorScatterUpdate->name()) {
|
||||
return CheckAndConvertUtils::CheckTensorTypeSame(type_dict, common_valid_types_with_bool, prim_name);
|
||||
}
|
||||
return CheckAndConvertUtils::CheckTensorTypeSame(type_dict, common_valid_types, prim_name);
|
||||
}
|
||||
} // namespace
|
||||
|
|
|
|||
Loading…
Reference in New Issue