mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-29 20:14:01 +00:00
Convert com.microsoft::ATen into org.pytorch.aten::ATen onnx op (#10060)
Signed-off-by: Thiago Crepaldi <thiago.crepaldi@microsoft.com>
This commit is contained in:
parent
e47434ea12
commit
e788cc2a23
24 changed files with 115 additions and 111 deletions
|
|
@ -13,6 +13,7 @@ constexpr const char* kOnnxDomain = "";
|
|||
constexpr const char* kOnnxDomainAlias = "ai.onnx";
|
||||
constexpr const char* kMLDomain = "ai.onnx.ml";
|
||||
constexpr const char* kMSDomain = "com.microsoft";
|
||||
constexpr const char* kPytorchAtenDomain = "org.pytorch.aten";
|
||||
constexpr const char* kMSExperimentalDomain = "com.microsoft.experimental";
|
||||
constexpr const char* kMSNchwcDomain = "com.microsoft.nchwc";
|
||||
constexpr const char* kMSInternalNHWCDomain = "com.ms.internal.nhwc";
|
||||
|
|
|
|||
|
|
@ -786,10 +786,10 @@ bool IsInputOnCpu(const Node& node, const KernelCreateInfo* p_kci, size_t index)
|
|||
}
|
||||
|
||||
#ifdef ENABLE_TRAINING
|
||||
if (node.GetExecutionProviderType() == kCudaExecutionProvider && node.OpType() == "ATenOp" && node.Domain() == kMSDomain) {
|
||||
if (node.GetExecutionProviderType() == kCudaExecutionProvider && node.OpType() == "ATen" && node.Domain() == kPytorchAtenDomain) {
|
||||
const auto& attrs = node.GetAttributes();
|
||||
ORT_ENFORCE(utils::HasString(attrs.at("name")));
|
||||
std::string op_name = attrs.at("name").s();
|
||||
ORT_ENFORCE(utils::HasString(attrs.at("operator")));
|
||||
std::string op_name = attrs.at("operator").s();
|
||||
std::string overload_name = "";
|
||||
if (attrs.find("overload_name") != attrs.end() && utils::HasString(attrs.at("overload_name"))) {
|
||||
overload_name = attrs.at("overload_name").s();
|
||||
|
|
|
|||
|
|
@ -77,7 +77,7 @@ struct ProviderHostCPUImpl : ProviderHostCPU {
|
|||
std::vector<int64_t>& split_sizes) override { return p->SplitBase::PrepareForCompute(input_shape, num_outputs, axis, before_dims, after_dims_including_split_axis, after_dims_excluding_split, split_sizes); }
|
||||
|
||||
// From cpu/tensor/concatbase.h (direct)
|
||||
Status ConcatBase__PrepareForCompute(const ConcatBase* p, OpKernelContext* ctx, const ConcatBase_InlinedTensorsVector& input_tensors, Prepare& prepare) override {
|
||||
Status ConcatBase__PrepareForCompute(const ConcatBase* p, OpKernelContext* ctx, const ConcatBase_InlinedTensorsVector& input_tensors, Prepare& prepare) override {
|
||||
return p->ConcatBase::PrepareForCompute(ctx, reinterpret_cast<const ConcatBase::InlinedTensorsVector&>(input_tensors), prepare); }
|
||||
|
||||
// GatherElements (direct)
|
||||
|
|
@ -176,7 +176,7 @@ struct ProviderHostCPUImpl : ProviderHostCPU {
|
|||
#endif
|
||||
|
||||
#ifdef ENABLE_TRAINING
|
||||
Status ATenOp__Compute(const contrib::ATenOp* p, OpKernelContext* p_ctx) override { return p->ATenOp::Compute(p_ctx); }
|
||||
Status ATen__Compute(const contrib::ATen* p, OpKernelContext* p_ctx) override { return p->ATen::Compute(p_ctx); }
|
||||
void contrib__record_event_in_tensor(const Tensor& event_id_tensor) override { return contrib::record_event_in_tensor(event_id_tensor); }
|
||||
void contrib__wait_event_in_tensor(const Tensor& event_id_tensor) override { return contrib::wait_event_in_tensor(event_id_tensor); }
|
||||
Status contrib__Group__Compute(const contrib::Group* p, OpKernelContext* context) override { return p->Group::Compute(context); }
|
||||
|
|
@ -189,7 +189,7 @@ struct ProviderHostCPUImpl : ProviderHostCPU {
|
|||
|
||||
// From aten_op.h (direct)
|
||||
bool contrib__IsATenOperatorExecutorInitialized() override { return contrib::IsATenOperatorExecutorInitialized(); }
|
||||
Status contrib__ExecuteReduceSumATenOp(OpKernelContext* p_ctx, const gsl::span<const int64_t>& axes, bool keepdims) override { return contrib::ExecuteReduceSumATenOp(p_ctx, axes, keepdims); }
|
||||
Status contrib__ExecuteReduceSumATen(OpKernelContext* p_ctx, const gsl::span<const int64_t>& axes, bool keepdims) override { return contrib::ExecuteReduceSumATen(p_ctx, axes, keepdims); }
|
||||
#endif
|
||||
#endif
|
||||
};
|
||||
|
|
|
|||
|
|
@ -137,7 +137,7 @@ struct ProviderHostCPU {
|
|||
#endif
|
||||
|
||||
#ifdef ENABLE_TRAINING
|
||||
virtual Status ATenOp__Compute(const contrib::ATenOp* p, OpKernelContext* p_ctx) = 0;
|
||||
virtual Status ATen__Compute(const contrib::ATen* p, OpKernelContext* p_ctx) = 0;
|
||||
virtual void contrib__record_event_in_tensor(const Tensor& event_id_tensor) = 0;
|
||||
virtual void contrib__wait_event_in_tensor(const Tensor& event_id_tensor) = 0;
|
||||
virtual Status contrib__Group__Compute(const contrib::Group* p, OpKernelContext* context) = 0;
|
||||
|
|
@ -150,7 +150,7 @@ struct ProviderHostCPU {
|
|||
|
||||
// From aten_op.h
|
||||
virtual bool contrib__IsATenOperatorExecutorInitialized() = 0;
|
||||
virtual Status contrib__ExecuteReduceSumATenOp(OpKernelContext* p_ctx, const gsl::span<const int64_t>& axes, bool keepdims) = 0;
|
||||
virtual Status contrib__ExecuteReduceSumATen(OpKernelContext* p_ctx, const gsl::span<const int64_t>& axes, bool keepdims) = 0;
|
||||
#endif
|
||||
#endif
|
||||
};
|
||||
|
|
@ -220,7 +220,7 @@ inline Status PrepareForTrainingCompute(const TensorShape& input_shape, int num_
|
|||
|
||||
// From aten_op.h
|
||||
inline bool IsATenOperatorExecutorInitialized() { return g_host_cpu.contrib__IsATenOperatorExecutorInitialized(); }
|
||||
inline Status ExecuteReduceSumATenOp(OpKernelContext* p_ctx, const gsl::span<const int64_t>& axes, bool keepdims) { return g_host_cpu.contrib__ExecuteReduceSumATenOp(p_ctx, axes, keepdims); }
|
||||
inline Status ExecuteReduceSumATen(OpKernelContext* p_ctx, const gsl::span<const int64_t>& axes, bool keepdims) { return g_host_cpu.contrib__ExecuteReduceSumATen(p_ctx, axes, keepdims); }
|
||||
} // namespace contrib
|
||||
#endif // ENABLE_TRAINING
|
||||
#endif // USE_CUDA || USE_ROCM
|
||||
|
|
|
|||
|
|
@ -739,7 +739,7 @@ Status ReduceKernel<allow_multi_axes>::ComputeImpl(OpKernelContext* ctx, cudnnRe
|
|||
}
|
||||
|
||||
#ifdef ENABLE_TRAINING
|
||||
// Use ATenOp for ReduceSum if possible.
|
||||
// Use ATen for ReduceSum if possible.
|
||||
const TensorShape& input_shape = X->Shape();
|
||||
if (contrib::IsATenOperatorExecutorInitialized() && cudnn_reduce_op == CUDNN_REDUCE_TENSOR_ADD && !calculate_log_ &&
|
||||
!calculate_sqt_ && !log_sum_exp_ && input_shape.Size() > 0) {
|
||||
|
|
@ -747,7 +747,7 @@ Status ReduceKernel<allow_multi_axes>::ComputeImpl(OpKernelContext* ctx, cudnnRe
|
|||
axes.resize(input_shape.NumDimensions());
|
||||
std::iota(axes.begin(), axes.end(), 0);
|
||||
}
|
||||
ORT_RETURN_IF_ERROR(contrib::ExecuteReduceSumATenOp(ctx, axes, keepdims_));
|
||||
ORT_RETURN_IF_ERROR(contrib::ExecuteReduceSumATen(ctx, axes, keepdims_));
|
||||
return Status::OK();
|
||||
}
|
||||
#endif
|
||||
|
|
|
|||
|
|
@ -167,7 +167,7 @@ template <typename T>
|
|||
struct EinsumTypedComputeProcessor;
|
||||
|
||||
namespace contrib {
|
||||
class ATenOp;
|
||||
class ATen;
|
||||
class Group;
|
||||
class PassThrough;
|
||||
class YieldOp;
|
||||
|
|
@ -228,6 +228,7 @@ struct DeleteOnUnloadPtr {
|
|||
|
||||
constexpr const char* kOnnxDomain = "";
|
||||
constexpr const char* kMSDomain = "com.microsoft";
|
||||
constexpr const char* kPytorchAtenDomain = "org.pytorch.aten";
|
||||
constexpr const char* kNGraphDomain = "com.intel.ai";
|
||||
constexpr const char* kCudaExecutionProvider = "CUDAExecutionProvider";
|
||||
constexpr const char* kDnnlExecutionProvider = "DnnlExecutionProvider";
|
||||
|
|
|
|||
|
|
@ -63,7 +63,7 @@ void operator delete(void* p, size_t /*size*/) noexcept { return Provider_GetHos
|
|||
namespace onnxruntime {
|
||||
#if defined(_MSC_VER) && !defined(__clang__)
|
||||
#pragma warning(push)
|
||||
// "Global initializer calls a non-constexpr function."
|
||||
// "Global initializer calls a non-constexpr function."
|
||||
#pragma warning(disable : 26426)
|
||||
#endif
|
||||
ProviderHost* g_host = Provider_GetHost();
|
||||
|
|
@ -577,7 +577,7 @@ Status Scan<9>::SetupSubgraphExecutionInfo(const SessionState& session_state, co
|
|||
|
||||
#ifdef ENABLE_TRAINING
|
||||
namespace contrib {
|
||||
Status ATenOp::Compute(OpKernelContext* p_ctx) const { return g_host_cpu.ATenOp__Compute(this, p_ctx); }
|
||||
Status ATen::Compute(OpKernelContext* p_ctx) const { return g_host_cpu.ATen__Compute(this, p_ctx); }
|
||||
Status Group::Compute(OpKernelContext* context) const { return g_host_cpu.contrib__Group__Compute(this, context); }
|
||||
Status PassThrough::Compute(OpKernelContext* context) const { return g_host_cpu.contrib__PassThrough__Compute(this, context); }
|
||||
Status YieldOp::Compute(OpKernelContext* context) const { return g_host_cpu.contrib__YieldOp__Compute(this, context); }
|
||||
|
|
|
|||
|
|
@ -223,6 +223,7 @@ Status Environment::Initialize(std::unique_ptr<logging::LoggingManager> logging_
|
|||
domainToVersionRangeInstance.AddDomainToVersion(onnxruntime::kMSExperimentalDomain, 1, 1);
|
||||
domainToVersionRangeInstance.AddDomainToVersion(onnxruntime::kMSNchwcDomain, 1, 1);
|
||||
domainToVersionRangeInstance.AddDomainToVersion(onnxruntime::kMSInternalNHWCDomain, 1, 1);
|
||||
domainToVersionRangeInstance.AddDomainToVersion(onnxruntime::kPytorchAtenDomain, 1, 1);
|
||||
#ifdef USE_DML
|
||||
domainToVersionRangeInstance.AddDomainToVersion(onnxruntime::kMSDmlDomain, 1, 1);
|
||||
#endif
|
||||
|
|
@ -272,7 +273,7 @@ Status Environment::Initialize(std::unique_ptr<logging::LoggingManager> logging_
|
|||
all_types.insert(all_types.end(), all_tensor_types.begin(), all_tensor_types.end());
|
||||
all_types.insert(all_types.end(), all_sequence_types.begin(), all_sequence_types.end());
|
||||
all_types.emplace_back("seq(tensor(bfloat16))");
|
||||
all_types.erase(std::remove_if(all_types.begin(), all_types.end(),
|
||||
all_types.erase(std::remove_if(all_types.begin(), all_types.end(),
|
||||
[](const std::string& s) { return s.find("string") != std::string::npos; }), all_types.end());
|
||||
return all_types; }();
|
||||
|
||||
|
|
|
|||
|
|
@ -1898,9 +1898,10 @@ class SymbolicShapeInference:
|
|||
vi = self.known_vi_[node.output[0]]
|
||||
if len(vi.type.tensor_type.shape.dim) == 0:
|
||||
vi.type.tensor_type.elem_type = onnx.TensorProto.UNDEFINED
|
||||
elif node.op_type == 'ATenOp' and node.domain == 'com.microsoft':
|
||||
elif node.op_type == 'ATen' and node.domain == 'org.pytorch.aten':
|
||||
for attr in node.attribute:
|
||||
if attr.name == 'name':
|
||||
# TODO: Is overload_name needed?
|
||||
if attr.name == 'operator':
|
||||
aten_op_name = attr.s.decode('utf-8') if isinstance(attr.s, bytes) else attr.s
|
||||
if aten_op_name in self.aten_op_dispatcher_:
|
||||
known_aten_op = True
|
||||
|
|
|
|||
|
|
@ -267,4 +267,4 @@
|
|||
"ZeroGradient com.microsoft CPUExecutionProvider",
|
||||
3284255990062374928
|
||||
]
|
||||
]
|
||||
]
|
||||
|
|
|
|||
|
|
@ -139,7 +139,7 @@ NodeSet GradientGraphBuilder::BFSWithStopGradient(const std::unordered_set<std::
|
|||
for (auto edge_it = n->OutputEdgesBegin(); edge_it != n->OutputEdgesEnd(); ++edge_it) {
|
||||
const Node& node = edge_it->GetNode();
|
||||
|
||||
const std::unordered_set<size_t>* edges = GetStopGradientEdges(node);
|
||||
const std::unordered_set<size_t>* edges = GetStopGradientEdges(node);
|
||||
if (edges != nullptr && edges->count(edge_it->GetDstArgIndex())) {
|
||||
continue;
|
||||
}
|
||||
|
|
@ -163,7 +163,7 @@ NodeSet GradientGraphBuilder::ReverseBFSWithStopGradient(const NodeSet& nodes) c
|
|||
queue.pop_front();
|
||||
|
||||
for (auto edge_it = n->InputEdgesBegin(); edge_it != n->InputEdgesEnd(); ++edge_it) {
|
||||
const std::unordered_set<size_t>* edges = GetStopGradientEdges(*n);
|
||||
const std::unordered_set<size_t>* edges = GetStopGradientEdges(*n);
|
||||
if (edges != nullptr && edges->count(edge_it->GetDstArgIndex())) {
|
||||
LOGS(logger_, INFO) << "Skip building gradient for input_" << edge_it->GetDstArgIndex()
|
||||
<< " of node: " << n->Name();
|
||||
|
|
@ -203,7 +203,7 @@ Status GradientGraphBuilder::CheckNodeArgsReachable() const {
|
|||
const std::unordered_set<size_t>* GradientGraphBuilder::GetStopGradientEdges(const Node& node) const {
|
||||
std::string op_type = node.OpType();
|
||||
|
||||
if (op_type == "ATenOp") {
|
||||
if (op_type == "ATen") {
|
||||
std::string key = GetGradientDefinitionKeyByNode(node);
|
||||
return GradientDefinitionRegistry::Instance().GetStopGradientEdgesForNode(key);
|
||||
} else {
|
||||
|
|
|
|||
|
|
@ -64,7 +64,7 @@ void ComputeBroadcastBackwardAxes(
|
|||
B_dim = B_dims[j].dim_param();
|
||||
if (A_dim != B_dim) {
|
||||
LOGS_DEFAULT(INFO) << "Gradient building for node " << node_name << ": symbolic dimension expects to match. " <<
|
||||
"A_dims:" << ToString(A_dims) << ", B_dims:" << ToString(B_dims) <<
|
||||
"A_dims:" << ToString(A_dims) << ", B_dims:" << ToString(B_dims) <<
|
||||
" This is a relaxing case, and the kernel might run into problem later if A_dims and B_dims turns out not broadcastable.";
|
||||
}
|
||||
} else if (A_dims[i].has_dim_param() && B_dims[j].has_dim_value()) {
|
||||
|
|
@ -73,8 +73,8 @@ void ComputeBroadcastBackwardAxes(
|
|||
|
||||
if (B_dim != 1) {
|
||||
LOGS_DEFAULT(INFO) << "Gradient building for node " << node_name << ": symbolic broadcasting expects the B_dimension to be 1. " <<
|
||||
"A_dims:" << ToString(A_dims) << ", B_dims:" << ToString(B_dims) <<
|
||||
" This is a relaxing case, and the kernel might run into problem later if A_dims and B_dims turns out not broadcastable.";
|
||||
"A_dims:" << ToString(A_dims) << ", B_dims:" << ToString(B_dims) <<
|
||||
" This is a relaxing case, and the kernel might run into problem later if A_dims and B_dims turns out not broadcastable.";
|
||||
} else {
|
||||
if (B_axes) {
|
||||
B_axes->push_back(gsl::narrow_cast<int64_t>(k));
|
||||
|
|
@ -86,7 +86,7 @@ void ComputeBroadcastBackwardAxes(
|
|||
|
||||
if (A_dim != 1) {
|
||||
LOGS_DEFAULT(INFO) << "Gradient building for node " << node_name << ": symbolic broadcasting expects the A_dimension to be 1. " <<
|
||||
"A_dims:" << ToString(A_dims) << ", B_dims:" << ToString(B_dims) <<
|
||||
"A_dims:" << ToString(A_dims) << ", B_dims:" << ToString(B_dims) <<
|
||||
" This is a relaxing case, and the kernel might run into problem later if A_dims and B_dims turns out not broadcastable.";
|
||||
} else {
|
||||
if (A_axes) {
|
||||
|
|
@ -130,10 +130,10 @@ Status GetShape(const ArgDef& arg_def, std::vector<Dimension>& shape) {
|
|||
std::string GetGradientDefinitionKeyByNode(const Node& node) {
|
||||
std::string op_type = node.OpType();
|
||||
std::string key = node.Domain() + "::" + op_type;
|
||||
if (op_type == "ATenOp") {
|
||||
if (op_type == "ATen") {
|
||||
const auto& attrs = node.GetAttributes();
|
||||
ORT_ENFORCE(utils::HasString(attrs.at("name")));
|
||||
key = key + "::" + attrs.at("name").s();
|
||||
ORT_ENFORCE(utils::HasString(attrs.at("operator")));
|
||||
key = key + "::" + attrs.at("operator").s();
|
||||
std::string overload_name = "";
|
||||
if (attrs.find("overload_name") != attrs.end() && utils::HasString(attrs.at("overload_name"))) {
|
||||
overload_name = attrs.at("overload_name").s();
|
||||
|
|
|
|||
|
|
@ -2100,16 +2100,16 @@ Example 4:
|
|||
.Const("C_MinusHalf", -0.5f, elem_type)
|
||||
.Const("C_alpha", kAlpha, elem_type)
|
||||
.Add(R"(
|
||||
ErfArg = Mul (X, C_SqrtHalf)
|
||||
ErfTerm = Erf (ErfArg)
|
||||
PartialSum = Add (ErfTerm, C_One)
|
||||
HalfPartialSum = Mul (C_Half, PartialSum)
|
||||
AlphaX = Mul (X, C_alpha)
|
||||
MinusHalfX = Mul (C_MinusHalf, X)
|
||||
ExpArg = Mul (MinusHalfX, X)
|
||||
ExpTerm = Exp (ExpArg)
|
||||
Term3 = Mul (AlphaX, ExpTerm)
|
||||
FullSum = Add (HalfPartialSum, Term3)
|
||||
ErfArg = Mul (X, C_SqrtHalf)
|
||||
ErfTerm = Erf (ErfArg)
|
||||
PartialSum = Add (ErfTerm, C_One)
|
||||
HalfPartialSum = Mul (C_Half, PartialSum)
|
||||
AlphaX = Mul (X, C_alpha)
|
||||
MinusHalfX = Mul (C_MinusHalf, X)
|
||||
ExpArg = Mul (MinusHalfX, X)
|
||||
ExpTerm = Exp (ExpArg)
|
||||
Term3 = Mul (AlphaX, ExpTerm)
|
||||
FullSum = Add (HalfPartialSum, Term3)
|
||||
dX = Mul (dY, FullSum)
|
||||
)");
|
||||
|
||||
|
|
@ -2994,18 +2994,18 @@ Return true if all elements are true and false otherwise.
|
|||
});
|
||||
|
||||
#ifdef ENABLE_TRAINING
|
||||
ONNX_CONTRIB_OPERATOR_SCHEMA(ATenOp)
|
||||
.SetDomain(kMSDomain)
|
||||
ONNX_CONTRIB_OPERATOR_SCHEMA(ATen)
|
||||
.SetDomain(kPytorchAtenDomain)
|
||||
.SinceVersion(1)
|
||||
.SetSupportLevel(OpSchema::SupportType::EXPERIMENTAL)
|
||||
.SetDoc("ATenOp")
|
||||
.Input(0, "inputs", "ATenOp inputs.", "T", OpSchema::Variadic,
|
||||
.SetDoc("ATen")
|
||||
.Input(0, "inputs", "ATen Op inputs.", "T", OpSchema::Variadic,
|
||||
/*is_homogeneous*/ false,
|
||||
/*min_arity*/ 1)
|
||||
.Output(0, "outputs", "ATenOp outputs.", "T", OpSchema::Variadic,
|
||||
.Output(0, "outputs", "ATen Op outputs.", "T", OpSchema::Variadic,
|
||||
/*is_homogeneous*/ false,
|
||||
/*min_arity*/ 1)
|
||||
.Attr("name", "Name of ATen operator.", AttributeProto::STRING)
|
||||
.Attr("operator", "Name of ATen operator.", AttributeProto::STRING)
|
||||
.Attr("overload_name", "Overload name of ATen operator.", AttributeProto::STRING, false)
|
||||
.TypeConstraint("T", OpSchema::all_tensor_types(), "Allow inputs and outputs to be any kind of tensor.");
|
||||
#endif
|
||||
|
|
|
|||
|
|
@ -3,9 +3,9 @@
|
|||
# Licensed under the MIT License.
|
||||
# --------------------------------------------------------------------------
|
||||
|
||||
# Support registration of ATenOp's symbolic (PyTorch exporter overriding) and gradient definition.
|
||||
# Support registration of ATen op's symbolic (PyTorch exporter overriding) and gradient definition.
|
||||
|
||||
# Each gradient definition of ATenOp is a list of node definitions: [node_def1, node_def2, ...].
|
||||
# Each gradient definition of ATen op is a list of node definitions: [node_def1, node_def2, ...].
|
||||
# Each node definition is a tuple: (op_type, inputs, outputs, attributes), while attributes is optional.
|
||||
# 'op_type' is a string or a tuple of two strings with op_name and domain respectively.
|
||||
# If it's string type, then the domain is the default ONNX domain: ''.
|
||||
|
|
@ -82,70 +82,70 @@ def register_gradient(domain, name, *attributes):
|
|||
return gradient_wrapper
|
||||
|
||||
|
||||
# For ATenOp, we need to provide op_name and overload name.
|
||||
@register_gradient('com.microsoft', 'ATenOp', 'aten::embedding', '')
|
||||
# For ATen op, we need to provide op_name and overload name.
|
||||
@register_gradient('org.pytorch.aten', 'ATen', 'aten::embedding', '')
|
||||
def embedding_gradient():
|
||||
return [
|
||||
('Constant', [], ['Const_0'], {'value': {'value': 0, 'dtype': 'int', 'is_tensor': True}}),
|
||||
('Shape', ['I(0)'], ['Shape_X']),
|
||||
('Gather', ['Shape_X', 'Const_0'], ['Gather_X_0'], {'axis': {'value': 0, 'dtype': 'int'}}),
|
||||
(('ATenOp', 'com.microsoft'), ['GO(0)', 'I(1)', 'Gather_X_0', 'I(2)', 'I(3)', 'I(4)'], [
|
||||
'GI(0)'], {'name': {'value': 'aten::embedding_backward', 'dtype': 'string'}}),
|
||||
(('ATen', 'org.pytorch.aten'), ['GO(0)', 'I(1)', 'Gather_X_0', 'I(2)', 'I(3)', 'I(4)'], [
|
||||
'GI(0)'], {'operator': {'value': 'aten::embedding_backward', 'dtype': 'string'}}),
|
||||
]
|
||||
|
||||
@register_gradient('com.microsoft', 'ATenOp', 'aten::diagonal', '')
|
||||
@register_gradient('org.pytorch.aten', 'ATen', 'aten::diagonal', '')
|
||||
def diagonal_gradient():
|
||||
return [
|
||||
('Shape', ['I(0)'], ['Shape_X']),
|
||||
(('ATenOp', 'com.microsoft'), ['GO(0)', 'Shape_X', 'I(1)', 'I(2)', 'I(3)'], [
|
||||
'GI(0)'], {'name': {'value': 'aten::diagonal_backward', 'dtype': 'string'}}),
|
||||
(('ATen', 'org.pytorch.aten'), ['GO(0)', 'Shape_X', 'I(1)', 'I(2)', 'I(3)'], [
|
||||
'GI(0)'], {'operator': {'value': 'aten::diagonal_backward', 'dtype': 'string'}}),
|
||||
]
|
||||
|
||||
@register_gradient('com.microsoft', 'ATenOp', 'aten::max_pool2d_with_indices', '')
|
||||
@register_gradient('org.pytorch.aten', 'ATen', 'aten::max_pool2d_with_indices', '')
|
||||
def max_pool2d_gradient():
|
||||
return [
|
||||
(('ATenOp', 'com.microsoft'), ['GO(0)', 'I(0)', 'I(1)', 'I(2)', 'I(3)', 'I(4)', 'I(5)', 'O(1)'], [
|
||||
'GI(0)'], {'name': {'value': 'aten::max_pool2d_with_indices_backward', 'dtype': 'string'}}),
|
||||
(('ATen', 'org.pytorch.aten'), ['GO(0)', 'I(0)', 'I(1)', 'I(2)', 'I(3)', 'I(4)', 'I(5)', 'O(1)'], [
|
||||
'GI(0)'], {'operator': {'value': 'aten::max_pool2d_with_indices_backward', 'dtype': 'string'}}),
|
||||
]
|
||||
|
||||
|
||||
@register_gradient('com.microsoft', 'ATenOp', 'aten::unfold', '')
|
||||
@register_gradient('org.pytorch.aten', 'ATen', 'aten::unfold', '')
|
||||
def unfold_gradient():
|
||||
return [
|
||||
('Shape', ['I(0)'], ['Shape_X']),
|
||||
(('ATenOp', 'com.microsoft'), ['GO(0)', 'Shape_X', 'I(1)', 'I(2)', 'I(3)'], [
|
||||
'GI(0)'], {'name': {'value': 'aten::unfold_backward', 'dtype': 'string'}}),
|
||||
(('ATen', 'org.pytorch.aten'), ['GO(0)', 'Shape_X', 'I(1)', 'I(2)', 'I(3)'], [
|
||||
'GI(0)'], {'operator': {'value': 'aten::unfold_backward', 'dtype': 'string'}}),
|
||||
]
|
||||
|
||||
|
||||
@register_gradient('com.microsoft', 'ATenOp', 'aten::avg_pool2d', '')
|
||||
@register_gradient('org.pytorch.aten', 'ATen', 'aten::avg_pool2d', '')
|
||||
def avg_pool2d_gradient():
|
||||
return [
|
||||
(('ATenOp', 'com.microsoft'), ['GO(0)', 'I(0)', 'I(1)', 'I(2)', 'I(3)', 'I(4)', 'I(5)', 'I(6)'], [
|
||||
'GI(0)'], {'name': {'value': 'aten::avg_pool2d_backward', 'dtype': 'string'}}),
|
||||
(('ATen', 'org.pytorch.aten'), ['GO(0)', 'I(0)', 'I(1)', 'I(2)', 'I(3)', 'I(4)', 'I(5)', 'I(6)'], [
|
||||
'GI(0)'], {'operator': {'value': 'aten::avg_pool2d_backward', 'dtype': 'string'}}),
|
||||
]
|
||||
|
||||
|
||||
@register_gradient('com.microsoft', 'ATenOp', 'aten::_adaptive_avg_pool2d', '')
|
||||
@register_gradient('org.pytorch.aten', 'ATen', 'aten::_adaptive_avg_pool2d', '')
|
||||
def adaptive_avg_pool2d_gradient():
|
||||
return [
|
||||
(('ATenOp', 'com.microsoft'), ['GO(0)', 'I(0)'], [
|
||||
'GI(0)'], {'name': {'value': 'aten::_adaptive_avg_pool2d_backward', 'dtype': 'string'}}),
|
||||
(('ATen', 'org.pytorch.aten'), ['GO(0)', 'I(0)'], [
|
||||
'GI(0)'], {'operator': {'value': 'aten::_adaptive_avg_pool2d_backward', 'dtype': 'string'}}),
|
||||
]
|
||||
|
||||
CustomGradientRegistry.register_custom_stop_gradient_edges([0], 'com.microsoft', 'ATenOp', 'aten::argmax', '')
|
||||
CustomGradientRegistry.register_custom_stop_gradient_edges([0], 'com.microsoft', 'ATenOp', 'aten::multinomial', '')
|
||||
CustomGradientRegistry.register_custom_stop_gradient_edges([0], 'org.pytorch.aten', 'ATen', 'aten::argmax', '')
|
||||
CustomGradientRegistry.register_custom_stop_gradient_edges([0], 'org.pytorch.aten', 'ATen', 'aten::multinomial', '')
|
||||
|
||||
@register_gradient('com.microsoft', 'ATenOp', 'aten::binary_cross_entropy_with_logits', '')
|
||||
@register_gradient('org.pytorch.aten', 'ATen', 'aten::binary_cross_entropy_with_logits', '')
|
||||
def binary_cross_entropy_with_logits_gradient():
|
||||
return [
|
||||
(('ATenOp', 'com.microsoft'), ['GO(0)', 'I(0)', 'I(1)', 'I(2)', 'I(3)', 'I(4)'], [
|
||||
'GI(0)'], {'name': {'value': 'aten::binary_cross_entropy_with_logits_backward', 'dtype': 'string'}}),
|
||||
(('ATen', 'org.pytorch.aten'), ['GO(0)', 'I(0)', 'I(1)', 'I(2)', 'I(3)', 'I(4)'], [
|
||||
'GI(0)'], {'operator': {'value': 'aten::binary_cross_entropy_with_logits_backward', 'dtype': 'string'}}),
|
||||
]
|
||||
|
||||
@register_gradient('com.microsoft', 'ATenOp', 'aten::numpy_T', '')
|
||||
@register_gradient('org.pytorch.aten', 'ATen', 'aten::numpy_T', '')
|
||||
def numpy_T_gradient():
|
||||
return [
|
||||
(('ATenOp', 'com.microsoft'), ['GO(0)'], [
|
||||
'GI(0)'], {'name': {'value': 'aten::numpy_T', 'dtype': 'string'}}),
|
||||
(('ATen', 'org.pytorch.aten'), ['GO(0)'], [
|
||||
'GI(0)'], {'operator': {'value': 'aten::numpy_T', 'dtype': 'string'}}),
|
||||
]
|
||||
|
|
|
|||
|
|
@ -64,8 +64,8 @@ def nll_loss(g, self, target, weight, reduction, ignore_index):
|
|||
|
||||
@register_symbolic('embedding')
|
||||
def embedding(g, weight, indices, padding_idx, scale_grad_by_freq, sparse):
|
||||
output = g.op("com.microsoft::ATenOp", weight, indices, padding_idx, scale_grad_by_freq, sparse,
|
||||
name_s='aten::embedding')
|
||||
output = g.op("org.pytorch.aten::ATen", weight, indices, padding_idx, scale_grad_by_freq, sparse,
|
||||
operator_s='aten::embedding')
|
||||
indices_shape = _get_tensor_sizes(indices)
|
||||
if indices_shape is not None and hasattr(weight.type(), 'with_sizes'):
|
||||
output_type = weight.type().with_sizes(
|
||||
|
|
@ -75,21 +75,21 @@ def embedding(g, weight, indices, padding_idx, scale_grad_by_freq, sparse):
|
|||
|
||||
@register_symbolic('bitwise_or')
|
||||
def bitwise_or(g, self, other):
|
||||
return g.op("com.microsoft::ATenOp", self, other,
|
||||
name_s='aten::bitwise_or', overload_name_s='Tensor')
|
||||
return g.op("org.pytorch.aten::ATen", self, other,
|
||||
operator_s='aten::bitwise_or', overload_name_s='Tensor')
|
||||
|
||||
@register_symbolic('diagonal')
|
||||
def diagonal(g, self, offset, dim1, dim2):
|
||||
return g.op("com.microsoft::ATenOp", self, offset, dim1, dim2,
|
||||
name_s='aten::diagonal')
|
||||
return g.op("org.pytorch.aten::ATen", self, offset, dim1, dim2,
|
||||
operator_s='aten::diagonal')
|
||||
|
||||
|
||||
@register_symbolic('multinomial')
|
||||
def multinomial(g, self, num_samples, replacement=False, generator=None):
|
||||
if generator is not None and not sym_help._is_none(generator):
|
||||
raise RuntimeError("Unsupported: ONNX does not support generator for multinomial")
|
||||
return g.op("com.microsoft::ATenOp", self, num_samples, replacement, generator,
|
||||
name_s='aten::multinomial')
|
||||
return g.op("org.pytorch.aten::ATen", self, num_samples, replacement, generator,
|
||||
operator_s='aten::multinomial')
|
||||
|
||||
|
||||
@register_symbolic('max_pool2d')
|
||||
|
|
@ -97,18 +97,18 @@ def max_pool2d(g, self, kernel_size, stride, padding, dilation, ceil_mode):
|
|||
stride_val = sym_help._maybe_get_const(stride, 'is')
|
||||
if not stride_val:
|
||||
stride = kernel_size
|
||||
return g.op("com.microsoft::ATenOp", self, kernel_size, stride, padding, dilation, ceil_mode,
|
||||
name_s='aten::max_pool2d_with_indices', outputs=2)[0]
|
||||
return g.op("org.pytorch.aten::ATen", self, kernel_size, stride, padding, dilation, ceil_mode,
|
||||
operator_s='aten::max_pool2d_with_indices', outputs=2)[0]
|
||||
|
||||
|
||||
@register_symbolic('unfold')
|
||||
def unfold(g, input, dimension, size, step):
|
||||
return g.op("com.microsoft::ATenOp", input, dimension, size, step, name_s='aten::unfold')
|
||||
return g.op("org.pytorch.aten::ATen", input, dimension, size, step, operator_s='aten::unfold')
|
||||
|
||||
|
||||
@register_symbolic('argmax')
|
||||
def argmax(g, input, dim, keepdim):
|
||||
return g.op("com.microsoft::ATenOp", input, dim, keepdim, name_s='aten::argmax')
|
||||
return g.op("org.pytorch.aten::ATen", input, dim, keepdim, operator_s='aten::argmax')
|
||||
|
||||
|
||||
@register_symbolic('avg_pool2d')
|
||||
|
|
@ -116,13 +116,13 @@ def avg_pool2d(g, self, kernel_size, stride, padding, ceil_mode, count_include_p
|
|||
stride_val = sym_help._maybe_get_const(stride, 'is')
|
||||
if not stride_val:
|
||||
stride = kernel_size
|
||||
return g.op("com.microsoft::ATenOp", self, kernel_size, stride, padding, ceil_mode,
|
||||
count_include_pad, divisor_override, name_s='aten::avg_pool2d')
|
||||
return g.op("org.pytorch.aten::ATen", self, kernel_size, stride, padding, ceil_mode,
|
||||
count_include_pad, divisor_override, operator_s='aten::avg_pool2d')
|
||||
|
||||
|
||||
@register_symbolic('adaptive_avg_pool2d')
|
||||
def adaptive_avg_pool2d(g, self, output_size):
|
||||
return g.op("com.microsoft::ATenOp", self, output_size, name_s='aten::_adaptive_avg_pool2d')
|
||||
return g.op("org.pytorch.aten::ATen", self, output_size, operator_s='aten::_adaptive_avg_pool2d')
|
||||
|
||||
|
||||
@register_symbolic('binary_cross_entropy_with_logits')
|
||||
|
|
@ -131,8 +131,8 @@ def binary_cross_entropy_with_logits(g, self, target, weight, pos_weight, reduct
|
|||
# But current custom_gradient_registry doesn't support such None checking,
|
||||
# So doesn't support non-None weight for now.
|
||||
if weight is None or sym_help._is_none(weight):
|
||||
return g.op("com.microsoft::ATenOp", self, target, weight, pos_weight, reduction,
|
||||
name_s='aten::binary_cross_entropy_with_logits')
|
||||
return g.op("org.pytorch.aten::ATen", self, target, weight, pos_weight, reduction,
|
||||
operator_s='aten::binary_cross_entropy_with_logits')
|
||||
from torch.onnx.symbolic_opset12 import binary_cross_entropy_with_logits as bce
|
||||
return bce(g, self, target, weight, pos_weight, reduction)
|
||||
|
||||
|
|
|
|||
|
|
@ -179,7 +179,7 @@ class GraphExecutionManager(GraphExecutionInterface):
|
|||
# Re-export will be avoided if _skip_check is enabled.
|
||||
self._original_model_has_changed = False
|
||||
|
||||
# Load ATenOp executor extension.
|
||||
# Load ATen operator executor extension.
|
||||
load_aten_op_executor_cpp_extension()
|
||||
|
||||
def _get_torch_gpu_allocator_function_addresses(self):
|
||||
|
|
|
|||
|
|
@ -343,7 +343,7 @@ def reinitialize_graph_execution_manager(graph_execution_manager):
|
|||
# Re-define the torch allocator
|
||||
graph_execution_manager._get_torch_gpu_allocator_function_addresses()
|
||||
|
||||
# Load ATenOp executor extension.
|
||||
# Load ATen op executor extension.
|
||||
load_aten_op_executor_cpp_extension()
|
||||
|
||||
def reinitialize_training_manager(training_manager):
|
||||
|
|
|
|||
|
|
@ -10,11 +10,11 @@
|
|||
namespace onnxruntime {
|
||||
namespace contrib {
|
||||
|
||||
ONNX_OPERATOR_KERNEL_EX(ATenOp, kMSDomain, 1, kCpuExecutionProvider,
|
||||
ONNX_OPERATOR_KERNEL_EX(ATen, kPytorchAtenDomain, 1, kCpuExecutionProvider,
|
||||
KernelDefBuilder().TypeConstraint("T", DataTypeImpl::AllTensorAndSequenceTensorTypes()),
|
||||
ATenOp);
|
||||
ATen);
|
||||
|
||||
Status ATenOp::Compute(OpKernelContext* p_ctx) const {
|
||||
Status ATen::Compute(OpKernelContext* p_ctx) const {
|
||||
auto* p_ctx_internal = static_cast<OpKernelContextInternal*>(p_ctx);
|
||||
std::vector<DLManagedTensor*> dlpacks;
|
||||
for (int i = 0; i < p_ctx_internal->InputCount(); i++) {
|
||||
|
|
@ -39,7 +39,7 @@ bool IsATenOperatorExecutorInitialized() {
|
|||
return aten_ops::ATenOperatorExecutor::Instance().IsInitialized();
|
||||
}
|
||||
|
||||
Status ExecuteReduceSumATenOp(OpKernelContext* p_ctx, const gsl::span<const int64_t>& axes, bool keepdims) {
|
||||
Status ExecuteReduceSumATen(OpKernelContext* p_ctx, const gsl::span<const int64_t>& axes, bool keepdims) {
|
||||
ORT_ENFORCE(aten_ops::ATenOperatorExecutor::Instance().IsInitialized() && !axes.empty());
|
||||
std::vector<DLManagedTensor*> dlpacks;
|
||||
auto* p_ctx_internal = static_cast<OpKernelContextInternal*>(p_ctx);
|
||||
|
|
|
|||
|
|
@ -8,10 +8,10 @@
|
|||
namespace onnxruntime {
|
||||
namespace contrib {
|
||||
|
||||
class ATenOp : public OpKernel {
|
||||
class ATen : public OpKernel {
|
||||
public:
|
||||
ATenOp(const OpKernelInfo& info) : OpKernel(info) {
|
||||
ORT_THROW_IF_ERROR(info.GetAttr("name", &op_name_));
|
||||
ATen(const OpKernelInfo& info) : OpKernel(info) {
|
||||
ORT_THROW_IF_ERROR(info.GetAttr("operator", &op_name_));
|
||||
overload_name_ = info.GetAttrOrDefault<std::string>("overload_name", "");
|
||||
}
|
||||
|
||||
|
|
@ -23,7 +23,7 @@ class ATenOp : public OpKernel {
|
|||
};
|
||||
|
||||
bool IsATenOperatorExecutorInitialized();
|
||||
Status ExecuteReduceSumATenOp(OpKernelContext* p_ctx, const gsl::span<const int64_t>& axes, bool keepdims);
|
||||
Status ExecuteReduceSumATen(OpKernelContext* p_ctx, const gsl::span<const int64_t>& axes, bool keepdims);
|
||||
|
||||
} // namespace contrib
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -98,7 +98,7 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, Recv)
|
|||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, RecordEvent);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, WaitEvent);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, YieldOp);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, ATenOp);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kPytorchAtenDomain, 1, ATen);
|
||||
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SummaryScalar);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SummaryHistogram);
|
||||
|
|
@ -208,7 +208,7 @@ Status RegisterCpuTrainingKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, RecordEvent)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, WaitEvent)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, YieldOp)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, ATenOp)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kPytorchAtenDomain, 1, ATen)>,
|
||||
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SummaryScalar)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCpuExecutionProvider, kMSDomain, 1, SummaryHistogram)>,
|
||||
|
|
|
|||
|
|
@ -8,9 +8,9 @@
|
|||
namespace onnxruntime {
|
||||
namespace cuda {
|
||||
|
||||
ONNX_OPERATOR_KERNEL_EX(ATenOp, kMSDomain, 1, kCudaExecutionProvider,
|
||||
ONNX_OPERATOR_KERNEL_EX(ATen, kPytorchAtenDomain, 1, kCudaExecutionProvider,
|
||||
(*KernelDefBuilder::Create()).TypeConstraint("T", DataTypeImpl::AllTensorAndSequenceTensorTypes()),
|
||||
onnxruntime::contrib::ATenOp);
|
||||
onnxruntime::contrib::ATen);
|
||||
|
||||
} // namespace cuda
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -215,7 +215,7 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, Adas
|
|||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, RecordEvent);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, WaitEvent);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, YieldOp);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, ATenOp);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kPytorchAtenDomain, 1, ATen);
|
||||
|
||||
#ifdef ENABLE_TRAINING_TORCH_INTEROP
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, PythonOp);
|
||||
|
|
@ -433,7 +433,7 @@ Status RegisterCudaTrainingKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, RecordEvent)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, WaitEvent)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, YieldOp)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, ATenOp)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kPytorchAtenDomain, 1, ATen)>,
|
||||
|
||||
#ifdef ENABLE_TRAINING_TORCH_INTEROP
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kCudaExecutionProvider, kMSDomain, 1, PythonOp)>,
|
||||
|
|
|
|||
|
|
@ -191,7 +191,7 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, Adas
|
|||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, RecordEvent);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, WaitEvent);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, YieldOp);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, ATenOp);
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kPytorchAtenDomain, 1, ATen);
|
||||
|
||||
#ifdef ENABLE_TRAINING_TORCH_INTEROP
|
||||
class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, PythonOp);
|
||||
|
|
@ -384,7 +384,7 @@ Status RegisterRocmTrainingKernels(KernelRegistry& kernel_registry) {
|
|||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, RecordEvent)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, WaitEvent)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, YieldOp)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, ATenOp)>,
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kPytorchAtenDomain, 1, ATen)>,
|
||||
|
||||
#ifdef ENABLE_TRAINING_TORCH_INTEROP
|
||||
BuildKernelCreateInfo<ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, PythonOp)>,
|
||||
|
|
|
|||
|
|
@ -27,9 +27,9 @@ def process_file(onnx_file):
|
|||
memcpu_ops.append(f"{node.op_type} {node.name}")
|
||||
if node.op_type == "Cast":
|
||||
cast_ops.append(f"{node.name}")
|
||||
if node.op_type == "ATenOp":
|
||||
if node.op_type == "ATen":
|
||||
for attr in node.attribute:
|
||||
if attr.name == "name":
|
||||
if attr.name == "operator":
|
||||
aten_ops.append(f"{node.name}: {attr.s.decode('utf-8')}")
|
||||
if node.op_type == "PythonOp":
|
||||
for attr in node.attribute:
|
||||
|
|
@ -51,7 +51,7 @@ def process_file(onnx_file):
|
|||
msgs.append(f"Examine whether {node.name} should be fused with the leading {prev.name} op into BiasSoftmax node.")
|
||||
|
||||
if aten_ops:
|
||||
print("ATenOp found:")
|
||||
print("ATen op found:")
|
||||
for line in aten_ops:
|
||||
print(line)
|
||||
print(10 * '-')
|
||||
|
|
|
|||
Loading…
Reference in a new issue