Cherry-pick 4 commits to rel-1.15.0 branch (#16021)

### Description

Cherry-pick 4 commits to rel-1.15.0 branch

#14834 
#15727
#16010
#16011
This commit is contained in:
Changming Sun 2023-05-19 21:48:21 -07:00 committed by GitHub
parent 6cdf071a94
commit ade58cf789
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
31 changed files with 1210 additions and 511 deletions

View file

@ -1,6 +1,6 @@
<?xml version="1.0" encoding="utf-8"?>
<packages>
<package id="python" version="3.9.7" targetFramework="native" />
<package id="Microsoft.AI.DirectML" version="1.11.0" targetFramework="native" />
<package id="Microsoft.AI.DirectML" version="1.12.0" targetFramework="native" />
<package id="Microsoft.Windows.CppWinRT" version="2.0.201201.7" targetFramework="native" />
</packages>

View file

@ -1,6 +1,6 @@
<?xml version="1.0" encoding="utf-8"?>
<packages>
<package id="pythonx86" version="3.9.7" targetFramework="native" />
<package id="Microsoft.AI.DirectML" version="1.11.0" targetFramework="native" />
<package id="Microsoft.AI.DirectML" version="1.12.0" targetFramework="native" />
<package id="Microsoft.Windows.CppWinRT" version="2.0.201201.7" targetFramework="native" />
</packages>

View file

@ -226,6 +226,7 @@ option(onnxruntime_BUILD_CACHE "onnxruntime build with cache" OFF)
cmake_dependent_option(MSVC_Z7_OVERRIDE "replacing /Zi and /ZI with /Z7 when using MSVC with CCache" ON "onnxruntime_BUILD_CACHE; MSVC" OFF)
option(onnxruntime_USE_AZURE "Build with azure inferencing support" OFF)
option(onnxruntime_USE_LOCK_FREE_QUEUE "Build with lock-free task queue for threadpool." OFF)
# ENABLE_TRAINING includes all training functionality
# The following 2 entry points
@ -743,7 +744,9 @@ if (onnxruntime_USE_AZURE)
list(APPEND ORT_PROVIDER_CMAKE_FLAGS -Donnxruntime_USE_AZURE=1)
list(APPEND ONNXRUNTIME_PROVIDER_NAMES azure)
endif()
if (onnxruntime_USE_LOCK_FREE_QUEUE)
add_compile_definitions(USE_LOCK_FREE_QUEUE)
endif()
if (onnxruntime_ENABLE_LAZY_TENSOR)
# To support LazyTensor, ORT needs to call Python function from C/C++.

View file

@ -41,7 +41,7 @@ if (NOT onnxruntime_USE_CUSTOM_DIRECTML)
set(NUGET_CONFIG ${PROJECT_SOURCE_DIR}/../NuGet.config)
set(PACKAGES_CONFIG ${PROJECT_SOURCE_DIR}/../packages.config)
get_filename_component(PACKAGES_DIR ${CMAKE_CURRENT_BINARY_DIR}/../packages ABSOLUTE)
set(DML_PACKAGE_DIR ${PACKAGES_DIR}/Microsoft.AI.DirectML.1.11.0)
set(DML_PACKAGE_DIR ${PACKAGES_DIR}/Microsoft.AI.DirectML.1.12.0)
# Restore nuget packages, which will pull down the DirectML redist package.
add_custom_command(

View file

@ -1193,8 +1193,10 @@ Do not modify directly.*
|DequantizeLinear|*in* x:**T1**<br> *in* x_scale:**T2**<br> *in* x_zero_point:**T1**<br> *out* y:**T2**|1+|**T1** = tensor(float)<br/> **T2** = tensor(uint8)|
|EmbedLayerNormalization|*in* input_ids:**T1**<br> *in* segment_ids:**T1**<br> *in* word_embedding:**T**<br> *in* position_embedding:**T**<br> *in* segment_embedding:**T**<br> *in* gamma:**T**<br> *in* beta:**T**<br> *in* mask:**T1**<br> *in* position_ids:**T1**<br> *out* output:**T**<br> *out* mask_index:**T1**<br> *out* embedding_sum:**T**|1+|**T** = tensor(float), tensor(float16)|
|FusedMatMul|*in* A:**T**<br> *in* B:**T**<br> *out* Y:**T**|1+|**T** = tensor(float), tensor(float16)|
|FusedMatMulActivation|*in* A:**T**<br> *in* B:**T**<br> *out* Y:**T**|1+|**T** = tensor(float), tensor(float16)|
|Gelu|*in* X:**T**<br> *out* Y:**T**|1+|**T** = tensor(float), tensor(float16)|
|GroupNorm|*in* X:**T**<br> *in* gamma:**M**<br> *in* beta:**M**<br> *out* Y:**T**|1+|**M** = tensor(float), tensor(float16)<br/> **T** = tensor(float), tensor(float16)|
|MultiHeadAttention|*in* query:**T**<br> *in* key:**T**<br> *in* value:**T**<br> *in* bias:**T**<br> *in* key_padding_mask:**M**<br> *in* relative_position_bias:**T**<br> *in* past_key:**T**<br> *in* past_value:**T**<br> *out* output:**T**<br> *out* present_key:**T**<br> *out* present_value:**T**|1+|**M** = tensor(int32)<br/> **T** = tensor(float), tensor(float16)|
|NhwcConv|*in* X:**T**<br> *in* W:**T**<br> *in* B:**T**<br> *out* Y:**T**|1+|**T** = tensor(float), tensor(float16)|
|QLinearAdd|*in* A:**T**<br> *in* A_scale:**tensor(float)**<br> *in* A_zero_point:**T**<br> *in* B:**T**<br> *in* B_scale:**tensor(float)**<br> *in* B_zero_point:**T**<br> *in* C_scale:**tensor(float)**<br> *in* C_zero_point:**T**<br> *out* C:**T**|1+|**T** = tensor(int8), tensor(uint8)|
|QLinearSigmoid|*in* X:**T**<br> *in* X_scale:**tensor(float)**<br> *in* X_zero_point:**T**<br> *in* Y_scale:**tensor(float)**<br> *in* Y_zero_point:**T**<br> *out* Y:**T**|1+|**T** = tensor(int8), tensor(uint8)|

View file

@ -44,6 +44,7 @@
#include "core/common/inlined_containers_fwd.h"
#include "core/common/spin_pause.h"
#include "core/platform/ort_mutex.h"
#include "core/platform/ort_spin_lock.h"
#include "core/platform/Barrier.h"
// ORT thread pool overview
@ -449,7 +450,11 @@ class RunQueue {
// PushBack adds w at the end of the queue.
// If queue is full returns w, otherwise returns default-constructed Work.
Work PushBack(Work w) {
#ifdef USE_LOCK_FREE_QUEUE
std::lock_guard<OrtSpinLock> mtx(spin_lock_);
#else
std::lock_guard<OrtMutex> lock(mutex_);
#endif
unsigned back = back_.load(std::memory_order_relaxed);
Elem& e = array_[(back - 1) & kMask];
ElemState s = e.state.load(std::memory_order_relaxed);
@ -469,7 +474,11 @@ class RunQueue {
// with w_idx. Typically the tag will be a per-thread ID to distinguish work
// submitted from different threads.
PushResult PushBackWithTag(Work w, Tag tag, unsigned& w_idx) {
#ifdef USE_LOCK_FREE_QUEUE
std::lock_guard<OrtSpinLock> mtx(spin_lock_);
#else
std::lock_guard<OrtMutex> lock(mutex_);
#endif
unsigned back = back_.load(std::memory_order_relaxed);
w_idx = (back - 1) & kMask;
Elem& e = array_[w_idx];
@ -490,7 +499,11 @@ class RunQueue {
Work PopBack() {
if (Empty())
return Work();
#ifdef USE_LOCK_FREE_QUEUE
std::lock_guard<OrtSpinLock> mtx(spin_lock_);
#else
std::lock_guard<OrtMutex> lock(mutex_);
#endif
unsigned back;
Elem* e;
ElemState s;
@ -532,7 +545,11 @@ class RunQueue {
bool RevokeWithTag(Tag tag, unsigned w_idx) {
bool revoked = false;
#ifdef USE_LOCK_FREE_QUEUE
std::lock_guard<OrtSpinLock> mtx(spin_lock_);
#else
std::lock_guard<OrtMutex> lock(mutex_);
#endif
Elem& e = array_[w_idx];
ElemState s = e.state.load(std::memory_order_relaxed);
@ -604,7 +621,11 @@ class RunQueue {
Work w;
};
#ifdef USE_LOCK_FREE_QUEUE
OrtSpinLock spin_lock_;
#else
OrtMutex mutex_;
#endif
// Low log(kSize) + 1 bits in front_ and back_ contain rolling index of
// front/back, respectively. The remaining bits contain modification counters

View file

@ -0,0 +1,35 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#pragma once
#include "core/common/spin_pause.h"
#include <atomic>
namespace onnxruntime {
/*
OrtSpinLock implemented mutex semantic "lock-freely",
calling thread will not be put to sleep on blocked,
which reduces cpu usage on context switching.
*/
struct OrtSpinLock {
using LockState = enum { Locked = 0,
Unlocked };
void lock() noexcept {
LockState state = Unlocked;
while (!state_.compare_exchange_weak(state, Locked, std::memory_order_acq_rel, std::memory_order_relaxed)) {
state = Unlocked;
concurrency::SpinPause(); // pause and retry
}
}
bool try_lock() noexcept {
LockState state = Unlocked;
return state_.compare_exchange_weak(state, Locked, std::memory_order_acq_rel, std::memory_order_relaxed);
}
void unlock() noexcept {
state_.store(Unlocked, std::memory_order_release);
}
private:
std::atomic<LockState> state_{Unlocked};
};
} // namespace onnxruntime

View file

@ -164,6 +164,7 @@ Status Attention<T>::ComputeInternal(OpKernelContext* context) const {
has_memory_efficient_attention(sm, sizeof(T) == 2);
#else
constexpr bool use_memory_efficient_attention = false;
ORT_UNUSED_VARIABLE(is_mask_1d_key_seq_len_start);
#endif
cublasHandle_t cublas = GetCublasHandle(context);

View file

@ -182,6 +182,7 @@ InlinedVector<std::unique_ptr<GraphTransformer>> GenerateTransformers(
#ifndef DISABLE_CONTRIB_OPS
const InlinedHashSet<std::string_view> cpu_ep = {onnxruntime::kCpuExecutionProvider};
#endif
const InlinedHashSet<std::string_view> dml_ep = {onnxruntime::kDmlExecutionProvider};
switch (level) {
case TransformerLevel::Level1: {
// RewriteRule optimizations are the simplest (they generally remove unnecessary nodes and are cheap to run)
@ -308,6 +309,7 @@ InlinedVector<std::unique_ptr<GraphTransformer>> GenerateTransformers(
transformers.emplace_back(std::make_unique<QuickGeluFusion>(cpu_cuda_dml_rocm_eps));
transformers.emplace_back(std::make_unique<MatMulScaleFusion>(cpu_cuda_dml_rocm_eps));
transformers.emplace_back(std::make_unique<MatMulActivationFusion>(dml_ep));
// GeluApproximation has side effects which may change results. It needs to be manually enabled,
// or alternatively the model can be updated offline using a model conversion script

View file

@ -183,13 +183,13 @@ private:
static T RoundUpToMultiple(T value, T multiple)
{
static_assert(std::is_integral_v<T>);
T remainder = value % multiple;
if (remainder != 0)
{
value += multiple - remainder;
}
return value;
}
@ -231,4 +231,4 @@ private:
// allocated memory if the fixed stack array is exhausted.
FixedBucket m_fixed;
std::deque<DynamicBucket> m_dynamic;
};
};

View file

@ -1155,6 +1155,11 @@ struct OperatorDescTraits<DML_ACTIVATION_GELU_OPERATOR_DESC>
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_ACTIVATION_GELU;
};
template <>
struct OperatorDescTraits<DML_MULTIHEAD_ATTENTION_OPERATOR_DESC>
{
static constexpr DML_OPERATOR_TYPE Type = DML_OPERATOR_MULTIHEAD_ATTENTION;
};
template <DML_OPERATOR_TYPE Type>
struct OperatorTypeTraits
@ -2139,14 +2144,20 @@ struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_ACTIVATION_GELU>
using DescType = DML_ACTIVATION_GELU_OPERATOR_DESC;
};
template <>
struct OperatorTypeTraits<(DML_OPERATOR_TYPE)DML_OPERATOR_MULTIHEAD_ATTENTION>
{
using DescType = DML_MULTIHEAD_ATTENTION_OPERATOR_DESC;
};
// Calls a visitor functor, supplying an empty operator desc corresponding to the given DML_OPERATOR_TYPE as
// the first argument.
//
//
// For example:
// Visit(DML_OPERATOR_ELEMENT_WISE_IDENTITY, [](auto tag) {
// using T = decltype(tag); // T is one of the DML_*_OPERATOR_DESC structs
// });
//
//
#pragma warning(push)
#pragma warning(disable:4702)
template <typename Visitor, typename... Ts>
@ -2432,6 +2443,8 @@ auto OperatorTypeVisitor(DML_OPERATOR_TYPE type, Visitor&& visitor, Ts&&... args
return std::invoke(std::forward<Visitor>(visitor), DML_RESAMPLE_GRAD1_OPERATOR_DESC{}, std::forward<Ts>(args)...);
case DML_OPERATOR_DIAGONAL_MATRIX1:
return std::invoke(std::forward<Visitor>(visitor), DML_DIAGONAL_MATRIX1_OPERATOR_DESC{}, std::forward<Ts>(args)...);
case DML_OPERATOR_MULTIHEAD_ATTENTION:
return std::invoke(std::forward<Visitor>(visitor), DML_MULTIHEAD_ATTENTION_OPERATOR_DESC{}, std::forward<Ts>(args)...);
case DML_OPERATOR_ACTIVATION_ELU:
return std::invoke(std::forward<Visitor>(visitor), DML_ACTIVATION_ELU_OPERATOR_DESC{}, std::forward<Ts>(args)...);
case DML_OPERATOR_ACTIVATION_CELU:
@ -2633,6 +2646,7 @@ inline gsl::czstring ToString(DML_OPERATOR_TYPE value)
case DML_OPERATOR_RESAMPLE2: return "DML_OPERATOR_RESAMPLE2";
case DML_OPERATOR_RESAMPLE_GRAD1: return "DML_OPERATOR_RESAMPLE_GRAD1";
case DML_OPERATOR_DIAGONAL_MATRIX1: return "DML_OPERATOR_DIAGONAL_MATRIX1";
case DML_OPERATOR_MULTIHEAD_ATTENTION: return "DML_OPERATOR_MULTIHEAD_ATTENTION";
default:
assert(false);
return "<unknown>";

View file

@ -2302,6 +2302,35 @@ constexpr DML_OPERATOR_SCHEMA DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA{
DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA_FIELDS,
};
constexpr DML_SCHEMA_FIELD DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA_FIELDS[18] {
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "QueryTensor", true },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "KeyTensor", true },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "ValueTensor", true },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "StackedQueryKeyTensor", true },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "StackedKeyValueTensor", true },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "StackedQueryKeyValueTensor", true },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "BiasTensor", true },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "MaskTensor", true },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "RelativePositionBiasTensor", true },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "PastKeyTensor", true },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "PastValueTensor", true },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputPresentKeyTensor", true },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputPresentValueTensor", true },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_FLOAT, "Scale", false },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_FLOAT, "MaskFilterValue", false },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "HeadCount", false },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_ATTRIBUTE, DML_SCHEMA_FIELD_TYPE_UINT, "MaskType", false },
};
constexpr DML_OPERATOR_SCHEMA DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA {
"DML_OPERATOR_MULTIHEAD_ATTENTION",
DML_OPERATOR_MULTIHEAD_ATTENTION,
DML_SCHEMA_OPERATOR_SUPPORT_FLAG_NONE,
18,
DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA_FIELDS,
};
constexpr DML_SCHEMA_FIELD DML_ACTIVATION_ELU_OPERATOR_SCHEMA_FIELDS[3] {
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_INPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "InputTensor", false },
DML_SCHEMA_FIELD { DML_SCHEMA_FIELD_KIND_OUTPUT_TENSOR, DML_SCHEMA_FIELD_TYPE_TENSOR_DESC, "OutputTensor", false },

View file

@ -202,7 +202,7 @@ namespace dml
};
}
#if DMLX_USE_ABSEIL
#if DMLX_USE_ABSEIL
template <typename T>
using Optional = absl::optional<T>;
@ -231,7 +231,7 @@ namespace dml
#elif DMLX_USE_GSL
template <typename T>
using Span = gsl::span<T>;
#else
#else
template <typename T>
using Span = dml::detail::span<T>;
#endif
@ -245,11 +245,11 @@ namespace dml
#define DMLX_THROW(_hr) THROW_HR(_hr)
#else
#define DMLX_THROW_IF_FAILED(_hr) if (FAILED(_hr)) { throw std::runtime_error(#_hr); }
#define DMLX_THROW(_hr) throw std::runtime_error(#_hr);
#define DMLX_THROW(_hr) throw std::runtime_error(#_hr);
#endif
#else
#define DMLX_THROW_IF_FAILED(_hr) if (FAILED(_hr)) { std::abort(); }
#define DMLX_THROW(_hr) { std::abort(); }
#define DMLX_THROW(_hr) { std::abort(); }
#endif
class Graph;
@ -307,7 +307,7 @@ namespace dml
// (0, 2, ..., n, 1). This is often referred to as "NHWC" or "interleaved channel" layout. This is useful,
// for example, when applied to 2D Convolution to produce outputs in an NHWC layout (as opposed to NCHW, which
// is the DirectML default for 2D Convolution).
//
//
// Examples of the transposes produced by this policy:
// NCW -> NWC
// NCHW -> NHWC
@ -713,7 +713,7 @@ namespace dml
// Represents an activation to be fused with an existing operator. The meaning of param1 and param2 depend on the
// activation to be fused.
//
//
// For HARD_SIGMOID, LINEAR, PARAMETRIC_SOFTPLUS, and SCALED_TANH: param1 = Alpha and param2 = Beta
// For ELU, LEAKY_RELU, THRESHOLDED_RELU, and CELU: param1 = Alpha. param2 is unused.
// For SCALED_ELU, param1 = Alpha and param2 = Gamma.
@ -1858,13 +1858,13 @@ namespace dml
}
// Helper for setting parameters for the Convolution operator. Sample usage:
//
//
// auto conv = dml::ConvolutionBuilder(...)
// .StartPadding(...)
// .EndPadding(...)
// .Strides(...)
// .Build();
//
//
// Parameters left unspecified will be defaulted with the same values as dml::Convolution().
class ConvolutionBuilder
{
@ -2114,9 +2114,9 @@ namespace dml
return output;
}
//
//
// TODO: LpPooling
//
//
// ---------------------------------------------------------------------------------------------------------------
@ -2203,13 +2203,13 @@ namespace dml
}
// Helper for setting parameters for the MaxPooling operator. Sample usage:
//
//
// auto [out, outIndices] = dml::MaxPoolingBuilder(...)
// .StartPadding(...)
// .EndPadding(...)
// .OutputIndices(...)
// .Build();
//
//
// Parameters left unspecified will be defaulted with the same values as dml::MaxPooling().
class MaxPoolingBuilder
{
@ -2251,13 +2251,13 @@ namespace dml
// ---------------------------------------------------------------------------------------------------------------
//
//
// TODO: MaxUnpooling
//
//
//
//
// TODO: ROIPooling
//
//
inline Expression Slice(
Expression input,
@ -2683,7 +2683,7 @@ namespace dml
{
detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder();
TensorDesc inputTensor = input.Impl()->GetOutputDesc();
assert(inputTensor.sizes.size() == 4);
dml::TensorDesc::Dimensions outputSizes = {
@ -2691,7 +2691,7 @@ namespace dml
inputTensor.sizes[1] * blockSize * blockSize,
inputTensor.sizes[2] / blockSize,
inputTensor.sizes[3] / blockSize
};
};
TensorDesc outputTensor(inputTensor.dataType, outputSizes, builder->GetTensorPolicy());
@ -2715,7 +2715,7 @@ namespace dml
{
detail::GraphBuilder* builder = input.Impl()->GetGraphBuilder();
TensorDesc inputTensor = input.Impl()->GetOutputDesc();
assert(inputTensor.sizes.size() == 4);
dml::TensorDesc::Dimensions outputSizes = {
@ -2771,7 +2771,7 @@ namespace dml
struct TopKOutputs
{
Expression value;
Expression index;
Expression index;
};
inline TopKOutputs TopK(Expression input, uint32_t axis, uint32_t k, DML_AXIS_DIRECTION axisDirection)
@ -2909,14 +2909,14 @@ namespace dml
desc.VarianceTensor = varianceTensor.AsPtr<DML_TENSOR_DESC>();
desc.ScaleTensor = scaleTensor.AsPtr<DML_TENSOR_DESC>();
desc.Epsilon = epsilon;
desc.OutputGradientTensor = outputGradientTensor.AsPtr<DML_TENSOR_DESC>();
desc.OutputScaleGradientTensor = outputScaleGradientTensor.AsPtr<DML_TENSOR_DESC>();
desc.OutputBiasGradientTensor = outputBiasGradientTensor.AsPtr<DML_TENSOR_DESC>();
dml::detail::NodeOutput* const inputs[] = { input.Impl(), inputGradient.Impl(), mean.Impl(), variance.Impl(), scale.Impl() };
dml::detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_BATCH_NORMALIZATION_GRAD, &desc, inputs);
BatchNormalizationGradOutputs outputValues;
outputValues.gradient = builder->CreateNodeOutput(node, 0, *desc.OutputGradientTensor);
outputValues.scaleGradient = builder->CreateNodeOutput(node, 1, *desc.OutputScaleGradientTensor);
@ -2932,7 +2932,7 @@ namespace dml
{
Expression output;
Expression mean;
Expression variance;
Expression variance;
};
inline BatchNormalizationTrainingOutputs BatchNormalizationTraining(
@ -3005,14 +3005,14 @@ namespace dml
desc.VarianceTensor = varianceTensor.AsPtr<DML_TENSOR_DESC>();
desc.ScaleTensor = scaleTensor.AsPtr<DML_TENSOR_DESC>();
desc.Epsilon = epsilon;
desc.OutputGradientTensor = outputGradientTensor.AsPtr<DML_TENSOR_DESC>();
desc.OutputScaleGradientTensor = outputScaleGradientTensor.AsPtr<DML_TENSOR_DESC>();
desc.OutputBiasGradientTensor = outputBiasGradientTensor.AsPtr<DML_TENSOR_DESC>();
dml::detail::NodeOutput* const inputs[] = { input.Impl(), inputGradient.Impl(), mean.Impl(), variance.Impl(), scale.Impl() };
dml::detail::NodeID node = builder->CreateOperatorNode(DML_OPERATOR_BATCH_NORMALIZATION_TRAINING_GRAD, &desc, inputs);
BatchNormalizationGradOutputs outputValues;
outputValues.gradient = builder->CreateNodeOutput(node, 0, *desc.OutputGradientTensor);
outputValues.scaleGradient = builder->CreateNodeOutput(node, 1, *desc.OutputScaleGradientTensor);
@ -3099,17 +3099,17 @@ namespace dml
return output;
}
//
//
// TODO: LpNormalization
//
//
//
//
// TODO: RNN
//
//
//
//
// TODO: LSTM
//
//
enum class GRUOutputOptions
{
@ -3121,7 +3121,7 @@ namespace dml
struct GRUOutputs
{
Expression sequence;
Expression single;
Expression single;
};
inline GRUOutputs GRU(
@ -3230,7 +3230,7 @@ namespace dml
return { outputSequenceExpr, outputSingleExpr };
}
//
//
// TODO: DiagonalMatrix
//
@ -3442,33 +3442,33 @@ namespace dml
return output;
}
//
//
// TODO: MatrixMultiplyInteger
//
//
//
//
// TODO: QuantizedLinearMatrixMultiply
//
//
//
//
// TODO: ConvolutionInteger
//
//
//
//
// TODO: QuantizedLinearConvolution
//
//
//
//
// TODO: ReluGrad
//
//
//
//
// TODO: AveragePoolingGrad
//
//
//
//
// TODO: MaxPoolingGrad
//
//
struct RandomGeneratorOutputs
{
@ -3496,7 +3496,7 @@ namespace dml
// Input and output state have the same TensorDesc.
desc.OutputStateTensor = inputStateTensor.AsPtr<DML_TENSOR_DESC>();
}
RandomGeneratorOutputs out;
detail::NodeOutput* const inputs[] = { inputState.Impl() };
@ -3537,7 +3537,7 @@ namespace dml
desc.InputTensor = inputTensor.AsPtr<DML_TENSOR_DESC>();
desc.OutputCountTensor = outputCountTensor.AsPtr<DML_TENSOR_DESC>();
desc.OutputCoordinatesTensor = outputCoordinatesTensor.AsPtr<DML_TENSOR_DESC>();
NonZeroCoordinatesOutputs output;
detail::NodeOutput* const inputs[] = { input.Impl() };
@ -3640,17 +3640,17 @@ namespace dml
return output;
}
//
//
// TODO: AdamOptimizer
//
//
//
//
// TODO: Argmin
//
//
//
//
// TODO: Argmax
//
//
#if DML_TARGET_VERSION >= 0x4000
@ -3694,7 +3694,7 @@ namespace dml
desc.ROITensor = roiTensor.AsPtr<DML_TENSOR_DESC>();
desc.BatchIndicesTensor = batchIndicesTensor.AsPtr<DML_TENSOR_DESC>();
desc.OutputTensor = outputTensor.AsPtr<DML_TENSOR_DESC>();
desc.ReductionFunction = reductionFunction;
desc.ReductionFunction = reductionFunction;
desc.InterpolationMode = interpolationMode;
desc.SpatialScaleX = spatialScaleX;
desc.SpatialScaleY = spatialScaleY;
@ -3763,7 +3763,7 @@ namespace dml
outputGradientTensor = TensorDesc(inputGradientTensor.dataType, outputGradientSizes, builder->GetTensorPolicy());
}
TensorDesc outputROIGradientTensor = computeOutputROIGradient ? TensorDesc(roiTensor.dataType, roiTensor.sizes, builder->GetTensorPolicy()) : TensorDesc();
assert(!computeOutputROIGradient || outputROIGradientTensor.sizes == roiTensor.sizes);
@ -3774,7 +3774,7 @@ namespace dml
desc.BatchIndicesTensor = batchIndicesTensor.AsPtr<DML_TENSOR_DESC>();
desc.OutputGradientTensor = computeOutputGradient ? outputGradientTensor.AsPtr<DML_TENSOR_DESC>() : nullptr;
desc.OutputROIGradientTensor = computeOutputROIGradient ? outputROIGradientTensor.AsPtr<DML_TENSOR_DESC>() : nullptr;
desc.ReductionFunction = reductionFunction;
desc.ReductionFunction = reductionFunction;
desc.InterpolationMode = interpolationMode;
desc.SpatialScaleX = spatialScaleX;
desc.SpatialScaleY = spatialScaleY;

View file

@ -1418,6 +1418,29 @@ inline std::vector<OperatorField> GetFields(const DML_DIAGONAL_MATRIX1_OPERATOR_
OperatorField(&DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA.Fields[5], ToOperatorFieldType(static_cast<INT>(desc.DiagonalFillEnd))),
};
}
inline std::vector<OperatorField> GetFields(const DML_MULTIHEAD_ATTENTION_OPERATOR_DESC& desc)
{
return {
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[0], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.QueryTensor))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[1], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.KeyTensor))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[2], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.ValueTensor))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[3], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.StackedQueryKeyTensor))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[4], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.StackedKeyValueTensor))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[5], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.StackedQueryKeyValueTensor))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[6], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.BiasTensor))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[7], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.MaskTensor))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[8], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.RelativePositionBiasTensor))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[9], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.PastKeyTensor))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[10], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.PastValueTensor))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[11], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputTensor))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[12], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputPresentKeyTensor))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[13], ToOperatorFieldType(static_cast<const DML_TENSOR_DESC*>(desc.OutputPresentValueTensor))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[14], ToOperatorFieldType(static_cast<FLOAT>(desc.Scale))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[15], ToOperatorFieldType(static_cast<FLOAT>(desc.MaskFilterValue))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[16], ToOperatorFieldType(static_cast<UINT>(desc.HeadCount))),
OperatorField(&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA.Fields[17], ToOperatorFieldType(static_cast<UINT>(desc.MaskType))),
};
}
inline std::vector<OperatorField> GetFields(const DML_ACTIVATION_ELU_OPERATOR_DESC& desc)
{
return {
@ -1753,6 +1776,7 @@ inline const DML_OPERATOR_SCHEMA& GetSchema(DML_OPERATOR_TYPE operatorType)
case DML_OPERATOR_RESAMPLE2: return DML_RESAMPLE2_OPERATOR_SCHEMA;
case DML_OPERATOR_RESAMPLE_GRAD1: return DML_RESAMPLE_GRAD1_OPERATOR_SCHEMA;
case DML_OPERATOR_DIAGONAL_MATRIX1: return DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA;
case DML_OPERATOR_MULTIHEAD_ATTENTION: return DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA;
case DML_OPERATOR_ACTIVATION_ELU: return DML_ACTIVATION_ELU_OPERATOR_SCHEMA;
case DML_OPERATOR_ACTIVATION_CELU: return DML_ACTIVATION_CELU_OPERATOR_SCHEMA;
case DML_OPERATOR_ACTIVATION_HARDMAX: return DML_ACTIVATION_HARDMAX_OPERATOR_SCHEMA;
@ -2346,6 +2370,10 @@ inline AbstractOperatorDesc ConvertOperatorDesc(const DML_OPERATOR_DESC& opDesc)
return AbstractOperatorDesc(
&DML_DIAGONAL_MATRIX1_OPERATOR_SCHEMA,
GetFields(*static_cast<const DML_DIAGONAL_MATRIX1_OPERATOR_DESC*>(opDesc.Desc)));
case DML_OPERATOR_MULTIHEAD_ATTENTION:
return AbstractOperatorDesc(
&DML_MULTIHEAD_ATTENTION_OPERATOR_SCHEMA,
GetFields(*static_cast<const DML_MULTIHEAD_ATTENTION_OPERATOR_DESC*>(opDesc.Desc)));
case DML_OPERATOR_ACTIVATION_ELU:
return AbstractOperatorDesc(
&DML_ACTIVATION_ELU_OPERATOR_SCHEMA,

View file

@ -10,32 +10,20 @@ Abbreviations: B is batch_size, S is sequence_length, W is hidden_size
M A B C // M, A, B, and C are Inputs
| \ | /
Cast Gemm
| Gemm
| / | \
| / | \
| / | \
| Slice Slice Slice
Identity | | |
| | | |
| | | |
| Identity Identity Identity // The identities are used to transpose NCHW -> NHCW while
| | | | // keeping the GEMM strides as NCHW to better target metacommands
| | | |
| ----- |
----------- | |
\ | |
Gemm |
| |
| |
Softmax |
| /
| /
\ /
\ /
Gemm
|
ActivationLinear
|
Output // Final output
----------------- MHA -----
|
|
Output // Final output
This kernel creates a DML_GRAPH, as mentioned above.
For reference, refer to this Doc:
@ -49,51 +37,165 @@ public:
DmlOperatorAttention(const MLOperatorKernelCreationContext& kernelCreationContext)
: DmlOperator(kernelCreationContext)
{
ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetInputCount() >= 3);
enum DmlInputIndex : uint32_t
{
mhaQueryIndex,
mhaKeyIndex,
mhaValueIndex,
mhaStackedQueryKeyIndex,
mhaStackedKeyValueIndex,
mhaStackedQueryKeyValueIndex,
mhaBiasIndex,
mhaMaskIndex,
mhaRelativePositionBiasIndex,
mhaPastKeyIndex,
mhaPastValueIndex,
mhaInputCount,
};
enum InputIndex : uint32_t
{
inputIndex,
weightsIndex,
biasIndex,
maskIndex,
pastIndex,
relativePositionBiasIndex,
pastSequenceLengthIndex,
inputCount,
};
enum OutputIndex : uint32_t
{
outputIndex,
outputCount,
};
ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetInputCount() >= 2);
ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetOutputCount() >= 1);
const uint32_t dmlInputIndex = inputIndex;
const uint32_t dmlWeightsIndex = weightsIndex;
const uint32_t dmlBiasIndex = biasIndex;
const uint32_t dmlMaskIndex = maskIndex;
const uint32_t dmlRelativePositionBiasIndex = relativePositionBiasIndex;
const bool hasBias = kernelCreationContext.IsInputValid(biasIndex);
const bool hasMask = kernelCreationContext.IsInputValid(maskIndex);
const bool hasUnpaddedBounds = hasMask && kernelCreationContext.GetInputTensorDimensionCount(maskIndex) == 1;
const bool hasRelativePositionBias = kernelCreationContext.IsInputValid(relativePositionBiasIndex);
DmlOperator::Initialize(kernelCreationContext, std::nullopt, std::nullopt, std::nullopt, std::nullopt, 1);
std::vector<uint32_t> inputTensorShape = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(0);
std::vector<uint32_t> weightTensorShape = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(1);
std::vector<uint32_t> biasTensorShape = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(2);
std::vector<uint32_t> maskIndexTensorShape = kernelCreationContext.GetTensorShapeDescription().GetInputTensorShape(3);
ML_CHECK_VALID_ARGUMENT(inputTensorShape.size() == 3);
ML_CHECK_VALID_ARGUMENT(weightTensorShape.size() == 2);
ML_CHECK_VALID_ARGUMENT(biasTensorShape.size() == 1);
ML_CHECK_VALID_ARGUMENT(weightTensorShape[1] == biasTensorShape[0]);
ML_CHECK_VALID_ARGUMENT(biasTensorShape[0] % 3 == 0);
ML_CHECK_VALID_ARGUMENT(inputTensorShape[2] == weightTensorShape[0]);
// TODO: fix Attention kernel when maskIndexTensorShape is 1
// https://microsoft.visualstudio.com/OS/_workitems/edit/41893987
ML_CHECK_VALID_ARGUMENT(maskIndexTensorShape.size() > 1 && maskIndexTensorShape.size() <= 4);
const uint32_t batchSize = inputTensorShape[0];
const uint32_t sequenceLength = inputTensorShape[1];
const uint32_t hiddenSize = biasTensorShape[0] / 3;
const uint32_t numHeads = gsl::narrow_cast<uint32_t>(kernelCreationContext.GetAttribute<int64_t>(AttrName::NumHeads));
ML_CHECK_VALID_ARGUMENT(numHeads > 0); // to avoid process crash because of division by zero.
ML_CHECK_VALID_ARGUMENT(hiddenSize % numHeads == 0);
auto inputTensorShape = m_inputTensorDescs[dmlInputIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(inputTensorShape.size() == 3);
auto weightTensorShape = m_inputTensorDescs[dmlWeightsIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(weightTensorShape.size() == 2);
ML_CHECK_VALID_ARGUMENT(weightTensorShape[0] == inputTensorShape[2]);
const auto qkvHiddenSizes = kernelCreationContext.GetOptionalAttributeVectorInt32(AttrName::QkvHiddenSizes);
if (hasBias)
{
auto biasTensorShape = m_inputTensorDescs[dmlBiasIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(biasTensorShape.size() == 1);
ML_CHECK_VALID_ARGUMENT(weightTensorShape[1] == biasTensorShape[0]);
if (qkvHiddenSizes.empty())
{
ML_CHECK_VALID_ARGUMENT(biasTensorShape[0] % 3 == 0);
}
}
if (!qkvHiddenSizes.empty())
{
ML_CHECK_VALID_ARGUMENT(qkvHiddenSizes.size() == 3);
ML_CHECK_VALID_ARGUMENT(qkvHiddenSizes[0] == qkvHiddenSizes[1]);
}
else
{
ML_CHECK_VALID_ARGUMENT(weightTensorShape[1] % 3 == 0);
}
const uint32_t hiddenSize = qkvHiddenSizes.empty() ? weightTensorShape[1] / 3 : qkvHiddenSizes[0];
const uint32_t vHiddenSize = qkvHiddenSizes.empty() ? weightTensorShape[1] / 3 : qkvHiddenSizes[2];
const uint32_t headSize = hiddenSize / numHeads;
const uint32_t vHeadSize = vHiddenSize / numHeads;
const uint32_t batchSize = inputTensorShape[0];
const uint32_t sequenceLength = inputTensorShape[1];
uint32_t desiredWeightTensorShape[3] = {batchSize, weightTensorShape[0], 3 * hiddenSize};
uint32_t desiredBiasTensorShape[3] = {batchSize, sequenceLength, 3 * hiddenSize};
MLOperatorTensorDataType dataType = kernelCreationContext.GetInputEdgeDescription(0).tensorDataType;
uint32_t desiredWeightTensorShape[3] = {batchSize, weightTensorShape[0], hiddenSize + hiddenSize + vHiddenSize};
MLOperatorTensorDataType dataType = kernelCreationContext.GetInputEdgeDescription(inputIndex).tensorDataType;
// overwrite weightTensorDesc
m_inputTensorDescs[1] = TensorDesc::ConstructBroadcastedTensorDesc(dataType, desiredWeightTensorShape, weightTensorShape);
m_inputTensorDescs[dmlWeightsIndex] = TensorDesc::ConstructBroadcastedTensorDesc(dataType, desiredWeightTensorShape, weightTensorShape);
// overwrite biasTensorDesc
m_inputTensorDescs[2] = TensorDesc::ConstructBroadcastedTensorDesc(dataType, desiredBiasTensorShape, biasTensorShape);
uint32_t desiredBiasTensorShape[3] = {batchSize, sequenceLength, hiddenSize + hiddenSize + vHiddenSize};
if (hasBias)
{
auto biasTensorShape = m_inputTensorDescs[dmlBiasIndex].GetSizes();
m_inputTensorDescs[dmlBiasIndex] = TensorDesc::ConstructBroadcastedTensorDesc(dataType, desiredBiasTensorShape, biasTensorShape);
}
// overwrite maskIndexTensorDesc
uint32_t maskIndexDimensionCount = gsl::narrow_cast<uint32_t>(maskIndexTensorShape.size());
maskIndexTensorShape.insert(maskIndexTensorShape.begin() + 1, 4 - maskIndexDimensionCount, 1);
uint32_t desiredMaskIndexShape[4] {batchSize, numHeads, sequenceLength, sequenceLength};
MLOperatorTensorDataType maskTensorDataType = kernelCreationContext.GetInputEdgeDescription(3).tensorDataType;
m_inputTensorDescs[3] = TensorDesc::ConstructBroadcastedTensorDesc(maskTensorDataType, desiredMaskIndexShape, maskIndexTensorShape);
MLOperatorTensorDataType maskTensorDataType = MLOperatorTensorDataType::Undefined;
bool hasMaxSequenceMask = false;
DML_MULTIHEAD_ATTENTION_MASK_TYPE maskType = DML_MULTIHEAD_ATTENTION_MASK_TYPE_NONE;
if (hasMask)
{
if (hasUnpaddedBounds)
{
auto unpaddedKeyBoundsShape = m_inputTensorDescs[dmlMaskIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(unpaddedKeyBoundsShape.size() == 1);
// overwrite output tensor desc
uint32_t outputTensorShape[4] = {batchSize, sequenceLength, numHeads, headSize};
uint32_t outputTensorStrides[4] = {sequenceLength * numHeads * headSize, headSize, headSize * sequenceLength, 1};
m_outputTensorDescs[0] = TensorDesc(GetDmlDataTypeFromMlDataType(dataType), outputTensorShape, outputTensorStrides, 0);
const uint32_t batchGroupCount = unpaddedKeyBoundsShape[0] / batchSize;
ML_CHECK_VALID_ARGUMENT(batchGroupCount == 1 || batchGroupCount == 2);
uint32_t desiredShape[2] = {batchGroupCount, batchSize};
m_inputTensorDescs[dmlMaskIndex] = TensorDesc(
m_inputTensorDescs[dmlMaskIndex].GetDmlDataType(),
desiredShape);
maskType = batchGroupCount == 1
? DML_MULTIHEAD_ATTENTION_MASK_TYPE_KEY_SEQUENCE_LENGTH
: DML_MULTIHEAD_ATTENTION_MASK_TYPE_KEY_SEQUENCE_END_START;
}
else
{
auto maskIndexTensorShape = m_inputTensorDescs[dmlMaskIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(maskIndexTensorShape.size() > 1 && maskIndexTensorShape.size() <= 4);
maskType = DML_MULTIHEAD_ATTENTION_MASK_TYPE_BOOLEAN;
std::vector<uint32_t> reshapedMaskIndexTensorShape(maskIndexTensorShape.begin(), maskIndexTensorShape.end());
if (maskIndexTensorShape.size() == 4 && maskIndexTensorShape[2] != sequenceLength)
{
hasMaxSequenceMask = true;
ML_CHECK_VALID_ARGUMENT(maskIndexTensorShape[2] == maskIndexTensorShape[3]);
const uint32_t maxSequenceLength = maskIndexTensorShape[2];
uint32_t desiredMaskIndexShape[4] {batchSize, numHeads, maxSequenceLength, maxSequenceLength};
maskTensorDataType = kernelCreationContext.GetInputEdgeDescription(maskIndex).tensorDataType;
m_inputTensorDescs[dmlMaskIndex] = TensorDesc::ConstructBroadcastedTensorDesc(maskTensorDataType, desiredMaskIndexShape, reshapedMaskIndexTensorShape);
}
else
{
uint32_t maskIndexDimensionCount = gsl::narrow_cast<uint32_t>(maskIndexTensorShape.size());
reshapedMaskIndexTensorShape.insert(reshapedMaskIndexTensorShape.begin() + 1, 4 - maskIndexDimensionCount, 1);
uint32_t desiredMaskIndexShape[4] {batchSize, numHeads, sequenceLength, sequenceLength};
maskTensorDataType = kernelCreationContext.GetInputEdgeDescription(maskIndex).tensorDataType;
m_inputTensorDescs[dmlMaskIndex] = TensorDesc::ConstructBroadcastedTensorDesc(maskTensorDataType, desiredMaskIndexShape, reshapedMaskIndexTensorShape);
}
}
}
if (hasRelativePositionBias)
{
auto relativePositionBiasTensorShape = m_inputTensorDescs[dmlRelativePositionBiasIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(relativePositionBiasTensorShape.size() == 4);
ML_CHECK_VALID_ARGUMENT(relativePositionBiasTensorShape[0] == inputTensorShape[0]);
ML_CHECK_VALID_ARGUMENT(relativePositionBiasTensorShape[1] == numHeads);
ML_CHECK_VALID_ARGUMENT(relativePositionBiasTensorShape[2] == inputTensorShape[1]);
}
TensorDesc firstGemmOutputTensorDesc = TensorDesc::ConstructDefaultTensorDesc(dataType, desiredBiasTensorShape);
DML_TENSOR_DESC namedFirstGemmOutputTensorDesc = firstGemmOutputTensorDesc.GetDmlDesc();
@ -101,345 +203,335 @@ public:
std::vector<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
std::vector<DML_TENSOR_DESC> outputDescs = GetDmlOutputDescs();
DML_GEMM_OPERATOR_DESC xWeightOperatorDesc = {};
xWeightOperatorDesc.ATensor = &inputDescs[0];
xWeightOperatorDesc.BTensor = &inputDescs[1];
xWeightOperatorDesc.CTensor = &inputDescs[2];
xWeightOperatorDesc.OutputTensor = &namedFirstGemmOutputTensorDesc;
xWeightOperatorDesc.TransA = DML_MATRIX_TRANSFORM_NONE;
xWeightOperatorDesc.TransB = DML_MATRIX_TRANSFORM_NONE;
xWeightOperatorDesc.Alpha = 1.0f;
xWeightOperatorDesc.Beta = 1.0f;
xWeightOperatorDesc.FusedActivation = nullptr;
const DML_OPERATOR_DESC xWeightDesc {DML_OPERATOR_GEMM, &xWeightOperatorDesc};
DML_GEMM_OPERATOR_DESC gemmOperatorDesc = {};
gemmOperatorDesc.ATensor = &inputDescs[0];
gemmOperatorDesc.BTensor = &inputDescs[1];
std::array<uint32_t, 3> querySlicedTensorShape {batchSize, sequenceLength, hiddenSize};
TensorDesc querySlicedInputTensorDesc = TensorDesc::ConstructDefaultTensorDesc(dataType, querySlicedTensorShape);
DML_TENSOR_DESC namedQuerySlicedInputTensorDesc = querySlicedInputTensorDesc.GetDmlDesc();
std::array<uint32_t, 3> querySliceOffset = {0, 0, 0};
std::array<uint32_t, 3> keySliceOffset = {0, 0, hiddenSize};
std::array<uint32_t, 3> valueSliceOffset = {0, 0, 2 * hiddenSize};
std::array<uint32_t, 3> sliceSize = {batchSize, sequenceLength, hiddenSize};
std::array<int32_t, 3> strides = {1, 1, 1};
DML_SLICE1_OPERATOR_DESC querySlicedOperatorDesc = {};
querySlicedOperatorDesc.InputTensor = &namedFirstGemmOutputTensorDesc;
querySlicedOperatorDesc.OutputTensor = &namedQuerySlicedInputTensorDesc;
querySlicedOperatorDesc.DimensionCount = gsl::narrow_cast<uint32_t>(querySlicedTensorShape.size());
querySlicedOperatorDesc.InputWindowOffsets = querySliceOffset.data();
querySlicedOperatorDesc.InputWindowSizes = sliceSize.data();
querySlicedOperatorDesc.InputWindowStrides = strides.data();
const DML_OPERATOR_DESC querySlicedDesc = { DML_OPERATOR_SLICE1, &querySlicedOperatorDesc };
DML_SLICE1_OPERATOR_DESC keySlicedOperatorDesc = {};
keySlicedOperatorDesc.InputTensor = &namedFirstGemmOutputTensorDesc;
keySlicedOperatorDesc.OutputTensor = &namedQuerySlicedInputTensorDesc;
keySlicedOperatorDesc.DimensionCount = gsl::narrow_cast<uint32_t>(querySlicedTensorShape.size());
keySlicedOperatorDesc.InputWindowOffsets = keySliceOffset.data();
keySlicedOperatorDesc.InputWindowSizes = sliceSize.data();
keySlicedOperatorDesc.InputWindowStrides = strides.data();
const DML_OPERATOR_DESC keySlicedDesc = { DML_OPERATOR_SLICE1, &keySlicedOperatorDesc };
DML_SLICE1_OPERATOR_DESC valueSlicedOperatorDesc = {};
valueSlicedOperatorDesc.InputTensor = &namedFirstGemmOutputTensorDesc;
valueSlicedOperatorDesc.OutputTensor = &namedQuerySlicedInputTensorDesc;
valueSlicedOperatorDesc.DimensionCount = gsl::narrow_cast<uint32_t>(querySlicedTensorShape.size());
valueSlicedOperatorDesc.InputWindowOffsets = valueSliceOffset.data();
valueSlicedOperatorDesc.InputWindowSizes = sliceSize.data();
valueSlicedOperatorDesc.InputWindowStrides = strides.data();
const DML_OPERATOR_DESC valueSlicedDesc = { DML_OPERATOR_SLICE1, &valueSlicedOperatorDesc};
TensorDesc castedMaskIndexTensorDesc = TensorDesc::ConstructDefaultTensorDesc(dataType, desiredMaskIndexShape);
DML_TENSOR_DESC namedCastedMaskIndexTensorDesc = castedMaskIndexTensorDesc.GetDmlDesc();
DML_CAST_OPERATOR_DESC castMaskIndexOperatorDesc = {};
castMaskIndexOperatorDesc.InputTensor = &inputDescs[3];
castMaskIndexOperatorDesc.OutputTensor = &namedCastedMaskIndexTensorDesc;
const DML_OPERATOR_DESC castMaskIndexDesc = {DML_OPERATOR_CAST, &castMaskIndexOperatorDesc};
// The attention fusion in ORT expects this to be number to -10000.
// https://github.com/microsoft/onnxruntime/blob/main/onnxruntime/core/optimizer/attention_fusion_helper.h#L604
// The decomposed Attention performs: (M - 1.0) * -10000.0, where M is the 4th input of the Attention node.
// Above equation can be written as (M * -1000) + 10000.0
DML_SCALE_BIAS scaleBias = {};
scaleBias.Scale = -10000.0f;
scaleBias.Bias = 10000.0f;
DML_ELEMENT_WISE_IDENTITY_OPERATOR_DESC maskOperatorDesc = {};
maskOperatorDesc.InputTensor = &namedCastedMaskIndexTensorDesc;
maskOperatorDesc.OutputTensor = &namedCastedMaskIndexTensorDesc;
maskOperatorDesc.ScaleBias = &scaleBias;
const DML_OPERATOR_DESC maskDesc = {DML_OPERATOR_ELEMENT_WISE_IDENTITY, &maskOperatorDesc};
// original reshaped shape: [batchSize, seqenceLength, numHeads, headSize]
// transposed shape to [0, 2, 1, 3] -> [batchSize, numHeads, sequenceLength, headSize]
uint32_t reshapedTransposedQueryTensorShape[4] = {batchSize, numHeads, sequenceLength, headSize};
uint32_t reshapedTransposedQueryTensorStrides[4] = {sequenceLength * numHeads * headSize, headSize, numHeads * headSize, 1};
TensorDesc reshapedTransposedQueryTensorDesc = TensorDesc(
GetDmlDataTypeFromMlDataType(dataType),
reshapedTransposedQueryTensorShape,
reshapedTransposedQueryTensorStrides);
DML_TENSOR_DESC namedReshapedTransposedQueryTensorDesc = reshapedTransposedQueryTensorDesc.GetDmlDesc();
TensorDesc reshapedTransposedQueryOutputTensorDesc = TensorDesc(
GetDmlDataTypeFromMlDataType(dataType),
reshapedTransposedQueryTensorShape);
DML_TENSOR_DESC namedReshapedTransposedQueryOutputTensorDesc = reshapedTransposedQueryOutputTensorDesc.GetDmlDesc();
DML_ELEMENT_WISE_IDENTITY_OPERATOR_DESC transposedQueryOperatorDesc{};
transposedQueryOperatorDesc.InputTensor = &namedReshapedTransposedQueryTensorDesc;
transposedQueryOperatorDesc.OutputTensor = &namedReshapedTransposedQueryOutputTensorDesc;
const DML_OPERATOR_DESC transposedQueryDesc {DML_OPERATOR_ELEMENT_WISE_IDENTITY, &transposedQueryOperatorDesc};
uint32_t reshapedTransposedKeyTensorShape[4] = {batchSize, numHeads, headSize, sequenceLength};
uint32_t reshapedTransposedKeyTensorStrides[4] = {sequenceLength * numHeads * headSize, headSize, 1, numHeads * headSize};
TensorDesc reshapedTransposedKeyTensorDesc = TensorDesc(
GetDmlDataTypeFromMlDataType(dataType),
reshapedTransposedKeyTensorShape,
reshapedTransposedKeyTensorStrides);
DML_TENSOR_DESC namedReshapedTransposedKeyTensorDesc = reshapedTransposedKeyTensorDesc.GetDmlDesc();
TensorDesc reshapedTransposedKeyOutputTensorDesc = TensorDesc(
GetDmlDataTypeFromMlDataType(dataType),
reshapedTransposedKeyTensorShape);
DML_TENSOR_DESC namedReshapedTransposedKeyOutputTensorDesc = reshapedTransposedKeyOutputTensorDesc.GetDmlDesc();
DML_ELEMENT_WISE_IDENTITY_OPERATOR_DESC transposedKeyOperatorDesc{};
transposedKeyOperatorDesc.InputTensor = &namedReshapedTransposedKeyTensorDesc;
transposedKeyOperatorDesc.OutputTensor = &namedReshapedTransposedKeyOutputTensorDesc;
const DML_OPERATOR_DESC transposedKeyDesc {DML_OPERATOR_ELEMENT_WISE_IDENTITY, &transposedKeyOperatorDesc};
uint32_t queryKeyTensorShape[4] = {batchSize, numHeads, sequenceLength, sequenceLength};
TensorDesc queryKeyTensorDesc = TensorDesc::ConstructDefaultTensorDesc(dataType, queryKeyTensorShape);
DML_TENSOR_DESC namedQueryKeyTensorDesc = queryKeyTensorDesc.GetDmlDesc();
float alpha = static_cast<float>(1 / sqrt(headSize));
DML_GEMM_OPERATOR_DESC attentionScoreOperatorDesc = {};
attentionScoreOperatorDesc.ATensor = &namedReshapedTransposedQueryOutputTensorDesc;
attentionScoreOperatorDesc.BTensor = &namedReshapedTransposedKeyOutputTensorDesc;
attentionScoreOperatorDesc.CTensor = &namedCastedMaskIndexTensorDesc;
attentionScoreOperatorDesc.OutputTensor = &namedQueryKeyTensorDesc;
attentionScoreOperatorDesc.TransA = DML_MATRIX_TRANSFORM_NONE;
attentionScoreOperatorDesc.TransB = DML_MATRIX_TRANSFORM_NONE;
attentionScoreOperatorDesc.Alpha = alpha;
attentionScoreOperatorDesc.Beta = 0.0f;
attentionScoreOperatorDesc.FusedActivation = nullptr;
const DML_OPERATOR_DESC attentionScoreDesc {DML_OPERATOR_GEMM, &attentionScoreOperatorDesc};
std::array<uint32_t, 1> axes = {3};
DML_ACTIVATION_SOFTMAX1_OPERATOR_DESC softmaxOperatorDesc = {};
softmaxOperatorDesc.InputTensor = &namedQueryKeyTensorDesc;
softmaxOperatorDesc.OutputTensor = &namedQueryKeyTensorDesc;
softmaxOperatorDesc.AxisCount = gsl::narrow_cast<uint32_t>(axes.size());
softmaxOperatorDesc.Axes = axes.data();
const DML_OPERATOR_DESC softmaxDesc = {DML_OPERATOR_ACTIVATION_SOFTMAX1, &softmaxOperatorDesc};
uint32_t reshapedTransposedOutputTensorShape[4] {batchSize, numHeads, sequenceLength, headSize};
uint32_t reshapedTransposedOutputTensorStrides[4] {sequenceLength * numHeads * headSize, headSize * sequenceLength, headSize, 1};
TensorDesc reshapedTransposedOutputTensorDesc = TensorDesc(
GetDmlDataTypeFromMlDataType(dataType),
reshapedTransposedOutputTensorShape,
reshapedTransposedOutputTensorStrides,
0 // guaranteedBaseOffsetAlignment
);
DML_TENSOR_DESC namedReshapedTransposedOutputTensorDesc = reshapedTransposedOutputTensorDesc.GetDmlDesc();
DML_GEMM_OPERATOR_DESC attentionWeightOperatorDesc = {};
attentionWeightOperatorDesc.ATensor = &namedQueryKeyTensorDesc;
attentionWeightOperatorDesc.BTensor = &namedReshapedTransposedQueryOutputTensorDesc;
attentionWeightOperatorDesc.CTensor = nullptr;
attentionWeightOperatorDesc.OutputTensor = &namedReshapedTransposedOutputTensorDesc;
attentionWeightOperatorDesc.TransA = DML_MATRIX_TRANSFORM_NONE;
attentionWeightOperatorDesc.TransB = DML_MATRIX_TRANSFORM_NONE;
attentionWeightOperatorDesc.Alpha = 1.0f;
attentionWeightOperatorDesc.Beta = 0.0f;
attentionWeightOperatorDesc.FusedActivation = nullptr;
const DML_OPERATOR_DESC attentionWeightDesc {DML_OPERATOR_GEMM, &attentionWeightOperatorDesc};
TensorDesc transposedOutputTensorDesc = TensorDesc(
m_outputTensorDescs[0].GetDmlDataType(),
m_outputTensorDescs[0].GetSizes(),
std::nullopt,
0 // guaranteedBaseOffsetAlignment
);
DML_TENSOR_DESC namedTransposedOutputTensorDesc = transposedOutputTensorDesc.GetDmlDesc();
DML_ACTIVATION_LINEAR_OPERATOR_DESC outputOperatorDesc = {};
outputOperatorDesc.Alpha = 1.0f;
outputOperatorDesc.Beta = 0.0f;
outputOperatorDesc.InputTensor = &outputDescs[0];
outputOperatorDesc.OutputTensor = &namedTransposedOutputTensorDesc;
const DML_OPERATOR_DESC outputDesc {DML_OPERATOR_ACTIVATION_LINEAR, &outputOperatorDesc};
enum NodeIndex : uint32_t
if (hasBias)
{
xWeight,
querySlice,
keySlice,
valueSlice,
queryTranspose,
keyTranspose,
attentionScore,
softmax,
valueTranspose,
attentionWeight,
castMaskIndex,
mask,
output,
count,
gemmOperatorDesc.CTensor = &inputDescs[2];
}
gemmOperatorDesc.OutputTensor = &namedFirstGemmOutputTensorDesc;
gemmOperatorDesc.TransA = DML_MATRIX_TRANSFORM_NONE;
gemmOperatorDesc.TransB = DML_MATRIX_TRANSFORM_NONE;
gemmOperatorDesc.Alpha = 1.0f;
gemmOperatorDesc.Beta = 1.0f;
gemmOperatorDesc.FusedActivation = nullptr;
const DML_OPERATOR_DESC gemmDesc {DML_OPERATOR_GEMM, &gemmOperatorDesc};
std::array<uint32_t, 3> queryKeySlicedTensorShape {batchSize, sequenceLength, hiddenSize + hiddenSize};
TensorDesc queryKeySlicedInputTensorDesc = TensorDesc::ConstructDefaultTensorDesc(dataType, queryKeySlicedTensorShape);
DML_TENSOR_DESC namedQueryKeySlicedInputTensorDesc = queryKeySlicedInputTensorDesc.GetDmlDesc();
std::array<uint32_t, 3> valueSlicedTensorShape {batchSize, sequenceLength, vHiddenSize};
TensorDesc valueSlicedInputTensorDesc = TensorDesc::ConstructDefaultTensorDesc(dataType, valueSlicedTensorShape);
DML_TENSOR_DESC namedValueSlicedInputTensorDesc = valueSlicedInputTensorDesc.GetDmlDesc();
// Transpose slice QK from [batchSize, sequenceLength, 2, numHeads, headSize] to [batchSize, sequenceLength, numHeads, 2, headSize]
std::array<uint32_t, 5> queryKeyTransposedTensorShape {batchSize, sequenceLength, numHeads, 2, headSize};
std::array<uint32_t, 5> queryKeyTransposedStrides {
sequenceLength * numHeads * 2 * headSize,
numHeads * 2 * headSize,
headSize,
numHeads * headSize,
1,
};
TensorDesc queryKeyTransposedInputTensorDesc = TensorDesc(
m_inputTensorDescs[dmlInputIndex].GetDmlDataType(),
queryKeyTransposedTensorShape,
queryKeyTransposedStrides);
DML_TENSOR_DESC namedQueryKeyTransposedInputTensorDesc = queryKeyTransposedInputTensorDesc.GetDmlDesc();
TensorDesc queryKeyTransposedOutputTensorDesc = TensorDesc(
m_inputTensorDescs[dmlInputIndex].GetDmlDataType(),
queryKeyTransposedTensorShape);
DML_TENSOR_DESC namedQueryKeyTransposedOutputTensorDesc = queryKeyTransposedOutputTensorDesc.GetDmlDesc();
// Transpose QKV from [batchSize, sequenceLength, 3, numHeads, headSize] to [batchSize, sequenceLength, numHeads, 3, headSize]
std::array<uint32_t, 5> queryKeyValueTransposedTensorShape {batchSize, sequenceLength, numHeads, 3, headSize};
std::array<uint32_t, 5> queryKeyValueTransposedStrides {
sequenceLength * numHeads * 3 * headSize,
numHeads * 3 * headSize,
headSize,
numHeads * headSize,
1,
};
TensorDesc queryKeyValueTransposedInputTensorDesc = TensorDesc(
m_inputTensorDescs[dmlInputIndex].GetDmlDataType(),
queryKeyValueTransposedTensorShape,
queryKeyValueTransposedStrides);
DML_TENSOR_DESC namedQueryKeyValueTransposedInputTensorDesc = queryKeyValueTransposedInputTensorDesc.GetDmlDesc();
TensorDesc queryKeyValueTransposedOutputTensorDesc = TensorDesc(
m_inputTensorDescs[dmlInputIndex].GetDmlDataType(),
queryKeyValueTransposedTensorShape);
DML_TENSOR_DESC namedQueryKeyValueTransposedOutputTensorDesc = queryKeyValueTransposedOutputTensorDesc.GetDmlDesc();
std::array<uint32_t, 3> queryKeySliceOffset = {0, 0, 0};
std::array<uint32_t, 3> queryKeySliceSize = {batchSize, sequenceLength, hiddenSize + hiddenSize};
std::array<int32_t, 3> queryKeySliceStrides = {1, 1, 1};
std::array<uint32_t, 3> valueSliceOffset = {0, 0, 2 * hiddenSize};
std::array<uint32_t, 3> valueSliceSize = {batchSize, sequenceLength, vHiddenSize};
std::array<int32_t, 3> valueSliceStrides = {1, 1, 1};
const bool hasSlicedValue = hiddenSize != vHiddenSize;
// We need to slice the value tensor when its hidden size is different from the query and key
DML_SLICE1_OPERATOR_DESC queryKeySlicedOperatorDesc = {};
DML_SLICE1_OPERATOR_DESC valueSlicedOperatorDesc = {};
DML_ELEMENT_WISE_IDENTITY_OPERATOR_DESC transposeOperatorDesc = {};
if (hasSlicedValue)
{
queryKeySlicedOperatorDesc.InputTensor = &namedFirstGemmOutputTensorDesc;
queryKeySlicedOperatorDesc.OutputTensor = &namedQueryKeySlicedInputTensorDesc;
queryKeySlicedOperatorDesc.DimensionCount = gsl::narrow_cast<uint32_t>(queryKeySlicedTensorShape.size());
queryKeySlicedOperatorDesc.InputWindowOffsets = queryKeySliceOffset.data();
queryKeySlicedOperatorDesc.InputWindowSizes = queryKeySliceSize.data();
queryKeySlicedOperatorDesc.InputWindowStrides = queryKeySliceStrides.data();
valueSlicedOperatorDesc.InputTensor = &namedFirstGemmOutputTensorDesc;
valueSlicedOperatorDesc.OutputTensor = &namedValueSlicedInputTensorDesc;
valueSlicedOperatorDesc.DimensionCount = gsl::narrow_cast<uint32_t>(valueSlicedTensorShape.size());
valueSlicedOperatorDesc.InputWindowOffsets = valueSliceOffset.data();
valueSlicedOperatorDesc.InputWindowSizes = valueSliceSize.data();
valueSlicedOperatorDesc.InputWindowStrides = valueSliceStrides.data();
transposeOperatorDesc.InputTensor = &namedQueryKeyTransposedInputTensorDesc;
transposeOperatorDesc.OutputTensor = &namedQueryKeyTransposedOutputTensorDesc;
}
else
{
// When Q/K/V all have the same hidden size, we just have to transpose it before sending it to MHA
transposeOperatorDesc.InputTensor = &namedQueryKeyValueTransposedInputTensorDesc;
transposeOperatorDesc.OutputTensor = &namedQueryKeyValueTransposedOutputTensorDesc;
}
const DML_OPERATOR_DESC queryKeySlicedDesc = { DML_OPERATOR_SLICE1, &queryKeySlicedOperatorDesc};
const DML_OPERATOR_DESC valueSlicedDesc = { DML_OPERATOR_SLICE1, &valueSlicedOperatorDesc};
const DML_OPERATOR_DESC transposedDesc = { DML_OPERATOR_ELEMENT_WISE_IDENTITY, &transposeOperatorDesc};
std::array<uint32_t, 4> maskSliceOutputShape {batchSize, numHeads, sequenceLength, sequenceLength};
std::array<int32_t, 4> maskSliceStrides = {1, 1, 1, 1};
std::array<uint32_t, 4> maskSliceOffsets = {0, 0, 0, 0};
TensorDesc maskSliceOutputTensorDesc;
DML_TENSOR_DESC namedMaskSliceOutputTensorDesc;
DML_SLICE1_OPERATOR_DESC maskSlicedOperatorDesc = {};
if (hasMaxSequenceMask)
{
maskSliceOutputTensorDesc = TensorDesc::ConstructDefaultTensorDesc(maskTensorDataType, maskSliceOutputShape);
namedMaskSliceOutputTensorDesc = maskSliceOutputTensorDesc.GetDmlDesc();
maskSlicedOperatorDesc.InputTensor = &inputDescs[dmlMaskIndex];
maskSlicedOperatorDesc.OutputTensor = &namedMaskSliceOutputTensorDesc;
maskSlicedOperatorDesc.DimensionCount = gsl::narrow_cast<uint32_t>(maskSliceOutputShape.size());
maskSlicedOperatorDesc.InputWindowOffsets = maskSliceOffsets.data();
maskSlicedOperatorDesc.InputWindowSizes = maskSliceOutputShape.data();
maskSlicedOperatorDesc.InputWindowStrides = maskSliceStrides.data();
}
const DML_OPERATOR_DESC maskSlicedDesc = { DML_OPERATOR_SLICE1, &maskSlicedOperatorDesc};
DML_MULTIHEAD_ATTENTION_OPERATOR_DESC mhaOperatorDesc = {};
mhaOperatorDesc.ValueTensor = hasSlicedValue ? &namedValueSlicedInputTensorDesc : nullptr;
mhaOperatorDesc.StackedQueryKeyTensor = hasSlicedValue ? &namedQueryKeyTransposedOutputTensorDesc : nullptr;
mhaOperatorDesc.StackedQueryKeyValueTensor = hasSlicedValue ? nullptr : &namedQueryKeyValueTransposedOutputTensorDesc;
if (hasMaxSequenceMask)
{
mhaOperatorDesc.MaskTensor = &namedMaskSliceOutputTensorDesc;
}
else
{
mhaOperatorDesc.MaskTensor = hasMask ? &inputDescs[dmlMaskIndex] : nullptr;
}
mhaOperatorDesc.RelativePositionBiasTensor = hasRelativePositionBias ? &inputDescs[dmlRelativePositionBiasIndex] : nullptr;
mhaOperatorDesc.OutputTensor = &outputDescs[outputIndex];
mhaOperatorDesc.Scale = kernelCreationContext.GetOptionalAttribute<float>(AttrName::Scale, gsl::narrow_cast<float>(1.0f / std::sqrt(headSize)));
mhaOperatorDesc.MaskFilterValue = kernelCreationContext.GetOptionalAttribute<float>(AttrName::MaskFilterValue, -10'000.0f);
mhaOperatorDesc.HeadCount = numHeads;
mhaOperatorDesc.MaskType = maskType;
const DML_OPERATOR_DESC mhaDesc = { DML_OPERATOR_MULTIHEAD_ATTENTION, &mhaOperatorDesc };
// Construct the graph
std::vector<DML_INPUT_GRAPH_EDGE_DESC> inputEdges;
std::vector<DML_INTERMEDIATE_GRAPH_EDGE_DESC> intermediateEdges;
std::vector<DML_OUTPUT_GRAPH_EDGE_DESC> outputEdges;
std::vector<const DML_OPERATOR_DESC*> opDescs = {
&gemmDesc,
&mhaDesc,
};
uint32_t currentNodeIndex = 0;
const uint32_t gemmNodeIndex = currentNodeIndex++;
const uint32_t mhaNodeIndex = currentNodeIndex++;
uint32_t valueSliceNodeIndex = 0;
uint32_t queryKeySliceNodeIndex = 0;
uint32_t queryKeyTransposedNodeIndex = 0;
uint32_t queryKeyValueTransposedNodeIndex = 0;
if (hasSlicedValue)
{
opDescs.push_back(&queryKeySlicedDesc);
queryKeySliceNodeIndex = currentNodeIndex++;
opDescs.push_back(&valueSlicedDesc);
valueSliceNodeIndex = currentNodeIndex++;
opDescs.push_back(&transposedDesc);
queryKeyTransposedNodeIndex = currentNodeIndex++;
}
else
{
opDescs.push_back(&transposedDesc);
queryKeyValueTransposedNodeIndex = currentNodeIndex++;
}
uint32_t maskSliceNodeIndex = 0;
if (hasMaxSequenceMask)
{
opDescs.push_back(&maskSlicedDesc);
maskSliceNodeIndex = currentNodeIndex++;
}
DML_INPUT_GRAPH_EDGE_DESC inputToGemmEdge = {};
inputToGemmEdge.GraphInputIndex = dmlInputIndex;
inputToGemmEdge.ToNodeIndex = gemmNodeIndex;
inputToGemmEdge.ToNodeInputIndex = 0;
inputEdges.push_back(inputToGemmEdge);
DML_INPUT_GRAPH_EDGE_DESC weightToGemmEdge = {};
weightToGemmEdge.GraphInputIndex = dmlWeightsIndex;
weightToGemmEdge.ToNodeIndex = gemmNodeIndex;
weightToGemmEdge.ToNodeInputIndex = 1;
inputEdges.push_back(weightToGemmEdge);
if (hasBias)
{
DML_INPUT_GRAPH_EDGE_DESC biasToGemmEdge = {};
biasToGemmEdge.GraphInputIndex = dmlBiasIndex;
biasToGemmEdge.ToNodeIndex = gemmNodeIndex;
biasToGemmEdge.ToNodeInputIndex = 2;
inputEdges.push_back(biasToGemmEdge);
}
if (hasMask)
{
if (hasUnpaddedBounds)
{
DML_INPUT_GRAPH_EDGE_DESC maskToMhaEdge = {};
maskToMhaEdge.GraphInputIndex = dmlMaskIndex;
maskToMhaEdge.ToNodeIndex = mhaNodeIndex;
maskToMhaEdge.ToNodeInputIndex = mhaMaskIndex;
inputEdges.push_back(maskToMhaEdge);
}
else if (hasMaxSequenceMask)
{
DML_INPUT_GRAPH_EDGE_DESC maskToMaskSliceEdge = {};
maskToMaskSliceEdge.GraphInputIndex = dmlMaskIndex;
maskToMaskSliceEdge.ToNodeIndex = maskSliceNodeIndex;
maskToMaskSliceEdge.ToNodeInputIndex = 0;
inputEdges.push_back(maskToMaskSliceEdge);
DML_INTERMEDIATE_GRAPH_EDGE_DESC maskSliceToMhaEdge = {};
maskSliceToMhaEdge.FromNodeIndex = maskSliceNodeIndex;
maskSliceToMhaEdge.FromNodeOutputIndex = 0;
maskSliceToMhaEdge.ToNodeIndex = mhaNodeIndex;
maskSliceToMhaEdge.ToNodeInputIndex = mhaMaskIndex;
intermediateEdges.push_back(maskSliceToMhaEdge);
}
else
{
DML_INPUT_GRAPH_EDGE_DESC maskToMhaEdge = {};
maskToMhaEdge.GraphInputIndex = dmlMaskIndex;
maskToMhaEdge.ToNodeIndex = mhaNodeIndex;
maskToMhaEdge.ToNodeInputIndex = mhaMaskIndex;
inputEdges.push_back(maskToMhaEdge);
}
}
if (hasRelativePositionBias)
{
DML_INPUT_GRAPH_EDGE_DESC relativePositionBiasToMhaEdge = {};
relativePositionBiasToMhaEdge.GraphInputIndex = dmlRelativePositionBiasIndex;
relativePositionBiasToMhaEdge.ToNodeIndex = mhaNodeIndex;
relativePositionBiasToMhaEdge.ToNodeInputIndex = mhaRelativePositionBiasIndex;
inputEdges.push_back(relativePositionBiasToMhaEdge);
}
if (hasSlicedValue)
{
// We need to slice QK and V, and transpose QK
DML_INTERMEDIATE_GRAPH_EDGE_DESC gemmToQueryKeySliceEdge = {};
gemmToQueryKeySliceEdge.FromNodeIndex = gemmNodeIndex;
gemmToQueryKeySliceEdge.FromNodeOutputIndex = 0;
gemmToQueryKeySliceEdge.ToNodeIndex = queryKeySliceNodeIndex;
gemmToQueryKeySliceEdge.ToNodeInputIndex = 0;
intermediateEdges.push_back(gemmToQueryKeySliceEdge);
DML_INTERMEDIATE_GRAPH_EDGE_DESC queryKeySliceToTransposeEdge = {};
queryKeySliceToTransposeEdge.FromNodeIndex = queryKeySliceNodeIndex;
queryKeySliceToTransposeEdge.FromNodeOutputIndex = 0;
queryKeySliceToTransposeEdge.ToNodeIndex = queryKeyTransposedNodeIndex;
queryKeySliceToTransposeEdge.ToNodeInputIndex = 0;
intermediateEdges.push_back(queryKeySliceToTransposeEdge);
DML_INTERMEDIATE_GRAPH_EDGE_DESC queryKeyTransposedToMhaEdge = {};
queryKeyTransposedToMhaEdge.FromNodeIndex = queryKeyTransposedNodeIndex;
queryKeyTransposedToMhaEdge.FromNodeOutputIndex = 0;
queryKeyTransposedToMhaEdge.ToNodeIndex = mhaNodeIndex;
queryKeyTransposedToMhaEdge.ToNodeInputIndex = mhaStackedQueryKeyIndex;
intermediateEdges.push_back(queryKeyTransposedToMhaEdge);
DML_INTERMEDIATE_GRAPH_EDGE_DESC gemmToValueSliceEdge = {};
gemmToValueSliceEdge.FromNodeIndex = gemmNodeIndex;
gemmToValueSliceEdge.FromNodeOutputIndex = 0;
gemmToValueSliceEdge.ToNodeIndex = valueSliceNodeIndex;
gemmToValueSliceEdge.ToNodeInputIndex = 0;
intermediateEdges.push_back(gemmToValueSliceEdge);
DML_INTERMEDIATE_GRAPH_EDGE_DESC valueSliceToMhaEdge = {};
valueSliceToMhaEdge.FromNodeIndex = valueSliceNodeIndex;
valueSliceToMhaEdge.FromNodeOutputIndex = 0;
valueSliceToMhaEdge.ToNodeIndex = mhaNodeIndex;
valueSliceToMhaEdge.ToNodeInputIndex = mhaValueIndex;
intermediateEdges.push_back(valueSliceToMhaEdge);
}
else
{
DML_INTERMEDIATE_GRAPH_EDGE_DESC gemmToQueryKeyValueTransposeEdge = {};
gemmToQueryKeyValueTransposeEdge.FromNodeIndex = gemmNodeIndex;
gemmToQueryKeyValueTransposeEdge.FromNodeOutputIndex = 0;
gemmToQueryKeyValueTransposeEdge.ToNodeIndex = queryKeyValueTransposedNodeIndex;
gemmToQueryKeyValueTransposeEdge.ToNodeInputIndex = 0;
intermediateEdges.push_back(gemmToQueryKeyValueTransposeEdge);
// All we need to do here is transpose the stacked QKV tensor into something DML supports
DML_INTERMEDIATE_GRAPH_EDGE_DESC queryKeyValueTransposedToMhaEdge = {};
queryKeyValueTransposedToMhaEdge.FromNodeIndex = queryKeyValueTransposedNodeIndex;
queryKeyValueTransposedToMhaEdge.FromNodeOutputIndex = 0;
queryKeyValueTransposedToMhaEdge.ToNodeIndex = mhaNodeIndex;
queryKeyValueTransposedToMhaEdge.ToNodeInputIndex = mhaStackedQueryKeyValueIndex;
intermediateEdges.push_back(queryKeyValueTransposedToMhaEdge);
}
DML_OUTPUT_GRAPH_EDGE_DESC mhaToOutputEdge = {};
mhaToOutputEdge.FromNodeIndex = mhaNodeIndex;
mhaToOutputEdge.FromNodeOutputIndex = 0;
mhaToOutputEdge.GraphOutputIndex = 0;
outputEdges.push_back(mhaToOutputEdge);
MLOperatorGraphDesc operatorGraphDesc = {};
std::array<const DML_OPERATOR_DESC*, NodeIndex::count> opDescs = {
&xWeightDesc,
&querySlicedDesc,
&keySlicedDesc,
&valueSlicedDesc,
&transposedQueryDesc,
&transposedKeyDesc,
&attentionScoreDesc,
&softmaxDesc,
&transposedQueryDesc,
&attentionWeightDesc,
&castMaskIndexDesc,
&maskDesc,
&outputDesc
};
operatorGraphDesc.nodeCount = NodeIndex::count;
operatorGraphDesc.nodesAsOpDesc = opDescs.data();
// set input edges
std::pair<uint32_t, uint32_t> nodeToNodeInputIndex[4] {
{NodeIndex::xWeight, 0},
{NodeIndex::xWeight, 1},
{NodeIndex::xWeight, 2},
{NodeIndex::castMaskIndex, 0}
};
std::array<DML_INPUT_GRAPH_EDGE_DESC, 4> inputEdges;
for (uint32_t inputIndex = 0; inputIndex < inputEdges.size(); inputIndex++)
{
DML_INPUT_GRAPH_EDGE_DESC inputEdge = {};
inputEdge.GraphInputIndex = inputIndex;
inputEdge.ToNodeIndex = nodeToNodeInputIndex[inputIndex].first;
inputEdge.ToNodeInputIndex = nodeToNodeInputIndex[inputIndex].second;
inputEdges[inputIndex] = inputEdge;
}
operatorGraphDesc.inputEdgeCount = gsl::narrow_cast<uint32_t>(inputEdges.size());
operatorGraphDesc.inputEdges = inputEdges.data();
// set intermediate edges
std::vector<DML_INTERMEDIATE_GRAPH_EDGE_DESC> intermediateEdges;
DML_INTERMEDIATE_GRAPH_EDGE_DESC gemmToQuerySliceEdge = {};
gemmToQuerySliceEdge.FromNodeIndex = NodeIndex::xWeight;
gemmToQuerySliceEdge.FromNodeOutputIndex = 0;
gemmToQuerySliceEdge.ToNodeIndex = NodeIndex::querySlice;
gemmToQuerySliceEdge.ToNodeInputIndex = 0;
intermediateEdges.push_back(gemmToQuerySliceEdge);
DML_INTERMEDIATE_GRAPH_EDGE_DESC gemmToKeySliceEdge = {};
gemmToKeySliceEdge.FromNodeIndex = NodeIndex::xWeight;
gemmToKeySliceEdge.FromNodeOutputIndex = 0;
gemmToKeySliceEdge.ToNodeIndex = NodeIndex::keySlice;
gemmToKeySliceEdge.ToNodeInputIndex = 0;
intermediateEdges.push_back(gemmToKeySliceEdge);
DML_INTERMEDIATE_GRAPH_EDGE_DESC gemmToValueSliceEdge = {};
gemmToValueSliceEdge.FromNodeIndex = NodeIndex::xWeight;
gemmToValueSliceEdge.FromNodeOutputIndex = 0;
gemmToValueSliceEdge.ToNodeIndex = NodeIndex::valueSlice;
gemmToValueSliceEdge.ToNodeInputIndex = 0;
intermediateEdges.push_back(gemmToValueSliceEdge);
DML_INTERMEDIATE_GRAPH_EDGE_DESC querySliceToQueryTranspose = {};
querySliceToQueryTranspose.FromNodeIndex = NodeIndex::querySlice;
querySliceToQueryTranspose.FromNodeOutputIndex = 0;
querySliceToQueryTranspose.ToNodeIndex = NodeIndex::queryTranspose;
querySliceToQueryTranspose.ToNodeInputIndex = 0;
intermediateEdges.push_back(querySliceToQueryTranspose);
DML_INTERMEDIATE_GRAPH_EDGE_DESC keySliceToKeyTranspose = {};
keySliceToKeyTranspose.FromNodeIndex = NodeIndex::keySlice;
keySliceToKeyTranspose.FromNodeOutputIndex = 0;
keySliceToKeyTranspose.ToNodeIndex = NodeIndex::keyTranspose;
keySliceToKeyTranspose.ToNodeInputIndex = 0;
intermediateEdges.push_back(keySliceToKeyTranspose);
DML_INTERMEDIATE_GRAPH_EDGE_DESC queryTransposeToGemm = {};
queryTransposeToGemm.FromNodeIndex = NodeIndex::queryTranspose;
queryTransposeToGemm.FromNodeOutputIndex = 0;
queryTransposeToGemm.ToNodeIndex = NodeIndex::attentionScore;
queryTransposeToGemm.ToNodeInputIndex = 0;
intermediateEdges.push_back(queryTransposeToGemm);
DML_INTERMEDIATE_GRAPH_EDGE_DESC keyTransposeToGemm = {};
keyTransposeToGemm.FromNodeIndex = NodeIndex::keyTranspose;
keyTransposeToGemm.FromNodeOutputIndex = 0;
keyTransposeToGemm.ToNodeIndex = NodeIndex::attentionScore;
keyTransposeToGemm.ToNodeInputIndex = 1;
intermediateEdges.push_back(keyTransposeToGemm);
DML_INTERMEDIATE_GRAPH_EDGE_DESC castedMaskIndexToIdentity = {};
castedMaskIndexToIdentity.FromNodeIndex = NodeIndex::castMaskIndex;
castedMaskIndexToIdentity.FromNodeOutputIndex = 0;
castedMaskIndexToIdentity.ToNodeIndex = NodeIndex::mask;
castedMaskIndexToIdentity.ToNodeInputIndex = 0;
intermediateEdges.push_back(castedMaskIndexToIdentity);
DML_INTERMEDIATE_GRAPH_EDGE_DESC maskToGemm = {};
maskToGemm.FromNodeIndex = NodeIndex::mask;
maskToGemm.FromNodeOutputIndex = 0;
maskToGemm.ToNodeIndex = NodeIndex::attentionScore;
maskToGemm.ToNodeInputIndex = 2;
intermediateEdges.push_back(maskToGemm);
DML_INTERMEDIATE_GRAPH_EDGE_DESC attentionScoreToSoftmax = {};
attentionScoreToSoftmax.FromNodeIndex = NodeIndex::attentionScore;
attentionScoreToSoftmax.FromNodeOutputIndex = 0;
attentionScoreToSoftmax.ToNodeIndex = NodeIndex::softmax;
attentionScoreToSoftmax.ToNodeInputIndex = 0;
intermediateEdges.push_back(attentionScoreToSoftmax);
DML_INTERMEDIATE_GRAPH_EDGE_DESC softmaxToGemm = {};
softmaxToGemm.FromNodeIndex = NodeIndex::softmax;
softmaxToGemm.FromNodeOutputIndex = 0;
softmaxToGemm.ToNodeIndex = NodeIndex::attentionWeight;
softmaxToGemm.ToNodeInputIndex = 0;
intermediateEdges.push_back(softmaxToGemm);
DML_INTERMEDIATE_GRAPH_EDGE_DESC valueSliceToValueTranspose = {};
valueSliceToValueTranspose.FromNodeIndex = NodeIndex::valueSlice;
valueSliceToValueTranspose.FromNodeOutputIndex = 0;
valueSliceToValueTranspose.ToNodeIndex = NodeIndex::valueTranspose;
valueSliceToValueTranspose.ToNodeInputIndex = 0;
intermediateEdges.push_back(valueSliceToValueTranspose);
DML_INTERMEDIATE_GRAPH_EDGE_DESC valueTransposeToGemm = {};
valueTransposeToGemm.FromNodeIndex = NodeIndex::valueTranspose;
valueTransposeToGemm.FromNodeOutputIndex = 0;
valueTransposeToGemm.ToNodeIndex = NodeIndex::attentionWeight;
valueTransposeToGemm.ToNodeInputIndex = 1;
intermediateEdges.push_back(valueTransposeToGemm);
DML_INTERMEDIATE_GRAPH_EDGE_DESC gemmToIdentity = {};
gemmToIdentity.FromNodeIndex = NodeIndex::attentionWeight;
gemmToIdentity.FromNodeOutputIndex = 0;
gemmToIdentity.ToNodeIndex = NodeIndex::output;
gemmToIdentity.ToNodeInputIndex = 0;
intermediateEdges.push_back(gemmToIdentity);
operatorGraphDesc.intermediateEdgeCount = gsl::narrow_cast<uint32_t>(intermediateEdges.size());
operatorGraphDesc.intermediateEdges = intermediateEdges.data();
// set the output edges
std::array<DML_OUTPUT_GRAPH_EDGE_DESC, 1> outputEdges;
DML_OUTPUT_GRAPH_EDGE_DESC outputEdge = {};
outputEdge.FromNodeIndex = NodeIndex::output;
outputEdge.FromNodeOutputIndex = 0;
outputEdge.GraphOutputIndex = 0;
outputEdges[0] = outputEdge;
operatorGraphDesc.outputEdgeCount = gsl::narrow_cast<uint32_t>(outputEdges.size());
operatorGraphDesc.outputEdges = outputEdges.data();
operatorGraphDesc.nodeCount = gsl::narrow_cast<uint32_t>(opDescs.size());
operatorGraphDesc.nodesAsOpDesc = opDescs.data();
SetDmlOperatorGraphDesc(std::move(operatorGraphDesc), kernelCreationContext);
}
@ -448,32 +540,37 @@ public:
void CALLBACK QueryAttention(IMLOperatorSupportQueryContextPrivate* context, /*out*/ bool* isSupported)
{
*isSupported = false;
// Fall back to CPU if input 'past' and 'relative_position_bias' is present because there is no current use case for this.
// and it will make the implementation more complex.
// Also fall back to CPU if output 'present' is present for same reason as above.
if (context->GetInputCount() > 4 || context->GetOutputCount() > 1)
{
return;
}
// Checking input count alone is not sufficient to fallback to CPU if input 'past' and 'relative_position_bias' is present
// because input 'mask_index', 'past', and 'relative_position_bias' all are optional.
if (context->IsInputValid(4) || context->IsInputValid(5))
{
return;
}
// Fall back to CPU if attibute 'qkv_hidden_sizes' is present or
// if value of attribute 'unidirectional' is 1, because of same reason as above.
MLOperatorAttributes attributes(context);
if (attributes.HasAttribute(AttrName::QkvHiddenSizes, MLOperatorAttributeType::IntArray))
// `past` input tensor is not supported yet
if (context->IsInputValid(4))
{
return;
}
// `past_sequence_length` input tensor is not supported yet
if (context->IsInputValid(6))
{
return;
}
// `present` output tensor is not supported yet
if (context->IsOutputValid(1))
{
return;
}
// `unidirectional == 1` is not supported yet
MLOperatorAttributes attributes(context);
if (attributes.GetOptionalAttribute<int32_t>(AttrName::Unidirectional, 0) != 0)
{
return;
}
// `do_rotary == 1` is not supported yet
if (attributes.GetOptionalAttribute<int32_t>(AttrName::DoRotary, 0) != 0)
{
return;
}
*isSupported = true;
}

View file

@ -0,0 +1,281 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include "precomp.h"
namespace Dml
{
class DmlOperatorMultiHeadAttention : public DmlOperator
{
public:
DmlOperatorMultiHeadAttention(const MLOperatorKernelCreationContext& kernelCreationContext)
: DmlOperator(kernelCreationContext)
{
enum InputIndex : uint32_t
{
queryIndex,
keyIndex,
valueIndex,
biasIndex,
maskIndex,
relativePositionBiasIndex,
pastKeyIndex,
pastValueIndex,
inputCount,
};
enum DmlInputIndex : uint32_t
{
dmlQueryIndex,
dmlKeyIndex,
dmlValueIndex,
dmlStackedQueryKeyIndex,
dmlStackedKeyValueIndex,
dmlStackedQueryKeyValueIndex,
dmlBiasIndex,
dmlMaskIndex,
dmlRelativePositionBiasIndex,
dmlPastKeyIndex,
dmlPastValueIndex,
dmlInputCount,
};
enum OutputIndex : uint32_t
{
outputIndex,
outputPresentKeyIndex,
outputPresentValueIndex,
outputCount,
};
ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetInputCount() >= 1);
ML_CHECK_VALID_ARGUMENT(kernelCreationContext.GetOutputCount() >= 1);
const bool keyValueIsPast = kernelCreationContext.IsInputValid(keyIndex) && kernelCreationContext.GetInputTensorDimensionCount(keyIndex) == 4;
const bool hasValue = kernelCreationContext.IsInputValid(valueIndex) && !keyValueIsPast;
const bool hasBias = kernelCreationContext.IsInputValid(biasIndex);
const bool hasMask = kernelCreationContext.IsInputValid(maskIndex);
const bool hasRelativePositionBias = kernelCreationContext.IsInputValid(relativePositionBiasIndex);
const bool hasPastKey = keyValueIsPast || kernelCreationContext.IsInputValid(pastKeyIndex);
const bool hasPastValue = keyValueIsPast || kernelCreationContext.IsInputValid(pastValueIndex);
const bool hasPresentKeyOutput = kernelCreationContext.IsOutputValid(outputPresentKeyIndex);
const bool hasPresentValueOutput = kernelCreationContext.IsOutputValid(outputPresentValueIndex);
const bool stackedQkv = kernelCreationContext.GetInputTensorDimensionCount(queryIndex) == 5;
const bool stackedKv = kernelCreationContext.IsInputValid(keyIndex) && kernelCreationContext.GetInputTensorDimensionCount(keyIndex) == 5;
const bool hasKey = !stackedKv && !keyValueIsPast && kernelCreationContext.IsInputValid(keyIndex);
std::vector<std::optional<uint32_t>> inputIndices = {
stackedQkv ? std::nullopt : std::optional<uint32_t>(queryIndex),
hasKey ? std::optional<uint32_t>(keyIndex) : std::nullopt,
hasValue ? std::optional<uint32_t>(valueIndex) : std::nullopt,
std::nullopt,
stackedKv ? std::optional<uint32_t>(keyIndex) : std::nullopt,
stackedQkv ? std::optional<uint32_t>(queryIndex) : std::nullopt,
biasIndex,
hasMask ? std::optional<uint32_t>(maskIndex) : std::nullopt,
relativePositionBiasIndex,
keyValueIsPast ? keyIndex : pastKeyIndex,
keyValueIsPast ? valueIndex : pastValueIndex,
};
std::vector<std::optional<uint32_t>> outputIndices = {
outputIndex,
outputPresentKeyIndex,
outputPresentValueIndex,
};
DmlOperator::Initialize(kernelCreationContext, inputIndices, outputIndices, std::nullopt, std::nullopt, 1);
ML_CHECK_VALID_ARGUMENT(!stackedQkv || m_inputTensorDescs[dmlStackedQueryKeyValueIndex].GetDimensionCount() == 5);
ML_CHECK_VALID_ARGUMENT(stackedQkv || m_inputTensorDescs[dmlQueryIndex].GetDimensionCount() == 3);
ML_CHECK_VALID_ARGUMENT(!hasKey || m_inputTensorDescs[dmlKeyIndex].GetDimensionCount() == 3);
ML_CHECK_VALID_ARGUMENT(!hasValue || m_inputTensorDescs[dmlValueIndex].GetDimensionCount() == 3);
ML_CHECK_VALID_ARGUMENT(!hasPastKey || m_inputTensorDescs[dmlPastKeyIndex].GetDimensionCount() == 4);
ML_CHECK_VALID_ARGUMENT(!hasPastValue || m_inputTensorDescs[dmlPastValueIndex].GetDimensionCount() == 4);
const uint32_t batchSize = stackedQkv
? m_inputTensorDescs[dmlStackedQueryKeyValueIndex].GetSizes()[0]
: m_inputTensorDescs[dmlQueryIndex].GetSizes()[0];
const uint32_t numHeads = gsl::narrow_cast<uint32_t>(kernelCreationContext.GetAttribute<int64_t>(AttrName::NumHeads));
const uint32_t headSize = stackedQkv
? m_inputTensorDescs[dmlStackedQueryKeyValueIndex].GetSizes()[4]
: m_inputTensorDescs[dmlQueryIndex].GetSizes()[2] / numHeads;
const uint32_t sequenceLength = stackedQkv
? m_inputTensorDescs[dmlStackedQueryKeyValueIndex].GetSizes()[1]
: m_inputTensorDescs[dmlQueryIndex].GetSizes()[1];
uint32_t kvSequenceLength;
if (hasKey)
{
kvSequenceLength = m_inputTensorDescs[dmlKeyIndex].GetSizes()[1];
}
else if (stackedKv)
{
kvSequenceLength = m_inputTensorDescs[dmlStackedKeyValueIndex].GetSizes()[1];
}
else if (hasPastKey)
{
kvSequenceLength = m_inputTensorDescs[dmlPastKeyIndex].GetSizes()[2];
}
else
{
kvSequenceLength = sequenceLength;
}
const uint32_t hiddenSize = numHeads * headSize;
const uint32_t vHiddenSize = hasValue ? m_inputTensorDescs[dmlValueIndex].GetSizes()[2] : hiddenSize;
const uint32_t pastSequenceLength = hasPastKey ? m_inputTensorDescs[dmlPastKeyIndex].GetSizes()[2] : 0;
const uint32_t totalSequenceLength = kvSequenceLength + pastSequenceLength;
if (stackedQkv)
{
auto stackedQkvSizes = m_inputTensorDescs[dmlStackedQueryKeyValueIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(stackedQkvSizes[0] == batchSize);
ML_CHECK_VALID_ARGUMENT(stackedQkvSizes[1] == sequenceLength);
ML_CHECK_VALID_ARGUMENT(stackedQkvSizes[2] == numHeads);
ML_CHECK_VALID_ARGUMENT(stackedQkvSizes[3] == 3);
ML_CHECK_VALID_ARGUMENT(stackedQkvSizes[4] == headSize);
}
else
{
auto querySizes = m_inputTensorDescs[dmlQueryIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(querySizes[0] == batchSize);
ML_CHECK_VALID_ARGUMENT(querySizes[1] == sequenceLength);
ML_CHECK_VALID_ARGUMENT(querySizes[2] == hiddenSize);
}
if (hasKey)
{
ML_CHECK_VALID_ARGUMENT(m_inputTensorDescs[dmlKeyIndex].GetDimensionCount() == 3);
auto keySizes = m_inputTensorDescs[dmlKeyIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(keySizes[0] == batchSize);
ML_CHECK_VALID_ARGUMENT(keySizes[1] == kvSequenceLength);
ML_CHECK_VALID_ARGUMENT(keySizes[2] == hiddenSize);
}
if (hasValue)
{
auto valueSizes = m_inputTensorDescs[dmlValueIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(valueSizes[0] == batchSize);
ML_CHECK_VALID_ARGUMENT(valueSizes[1] == kvSequenceLength);
ML_CHECK_VALID_ARGUMENT(valueSizes[2] == vHiddenSize);
}
if (stackedKv)
{
ML_CHECK_VALID_ARGUMENT(m_inputTensorDescs[dmlStackedKeyValueIndex].GetDimensionCount() == 5);
auto stackedKvSizes = m_inputTensorDescs[dmlStackedKeyValueIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(stackedKvSizes[0] == batchSize);
ML_CHECK_VALID_ARGUMENT(stackedKvSizes[1] == kvSequenceLength);
ML_CHECK_VALID_ARGUMENT(stackedKvSizes[2] == numHeads);
ML_CHECK_VALID_ARGUMENT(stackedKvSizes[3] == 2);
ML_CHECK_VALID_ARGUMENT(stackedKvSizes[4] == headSize);
}
if (hasBias)
{
ML_CHECK_VALID_ARGUMENT(m_inputTensorDescs[dmlBiasIndex].GetDimensionCount() == 1);
ML_CHECK_VALID_ARGUMENT(m_inputTensorDescs[dmlBiasIndex].GetSizes()[0] == hiddenSize + hiddenSize + vHiddenSize);
}
DML_MULTIHEAD_ATTENTION_MASK_TYPE maskType = DML_MULTIHEAD_ATTENTION_MASK_TYPE_NONE;
if (hasMask)
{
if (kernelCreationContext.GetInputTensorDimensionCount(maskIndex) == 1)
{
const auto unpaddedKeyBoundsShape = m_inputTensorDescs[dmlMaskIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(unpaddedKeyBoundsShape.size() == 1);
ML_CHECK_VALID_ARGUMENT(unpaddedKeyBoundsShape[0] == batchSize || unpaddedKeyBoundsShape[0] == batchSize * 3 + 2);
maskType = unpaddedKeyBoundsShape[0] == batchSize
? DML_MULTIHEAD_ATTENTION_MASK_TYPE_KEY_SEQUENCE_LENGTH
: DML_MULTIHEAD_ATTENTION_MASK_TYPE_KEY_QUERY_SEQUENCE_LENGTH_START_END;
if (maskType == DML_MULTIHEAD_ATTENTION_MASK_TYPE_KEY_SEQUENCE_LENGTH)
{
uint32_t desiredShape[2] = {1, batchSize};
m_inputTensorDescs[dmlMaskIndex] = TensorDesc(
m_inputTensorDescs[dmlMaskIndex].GetDmlDataType(),
desiredShape);
}
}
else
{
const auto keyPaddingMaskTensorShape = m_inputTensorDescs[dmlMaskIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(keyPaddingMaskTensorShape.size() == 2);
ML_CHECK_VALID_ARGUMENT(keyPaddingMaskTensorShape[0] == batchSize);
ML_CHECK_VALID_ARGUMENT(keyPaddingMaskTensorShape[1] == kvSequenceLength);
const uint32_t actualShape[4] = {batchSize, 1, 1, kvSequenceLength};
const uint32_t desiredShape[4] = {batchSize, numHeads, sequenceLength, kvSequenceLength};
m_inputTensorDescs[dmlMaskIndex] = TensorDesc::ConstructBroadcastedTensorDesc(
m_inputTensorDescs[dmlMaskIndex].GetMlOperatorDataType(),
desiredShape,
actualShape);
maskType = DML_MULTIHEAD_ATTENTION_MASK_TYPE_BOOLEAN;
}
}
if (hasRelativePositionBias)
{
ML_CHECK_VALID_ARGUMENT(m_inputTensorDescs[dmlRelativePositionBiasIndex].GetDimensionCount() == 4);
auto relativePositionBiasSizes = m_inputTensorDescs[dmlRelativePositionBiasIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(relativePositionBiasSizes[0] == batchSize);
ML_CHECK_VALID_ARGUMENT(relativePositionBiasSizes[1] == numHeads);
ML_CHECK_VALID_ARGUMENT(relativePositionBiasSizes[2] == sequenceLength);
ML_CHECK_VALID_ARGUMENT(relativePositionBiasSizes[3] == totalSequenceLength);
}
if (hasPastKey)
{
auto pastKeySizes = m_inputTensorDescs[dmlPastKeyIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(pastKeySizes[0] == batchSize);
ML_CHECK_VALID_ARGUMENT(pastKeySizes[1] == numHeads);
ML_CHECK_VALID_ARGUMENT(pastKeySizes[2] == pastSequenceLength);
ML_CHECK_VALID_ARGUMENT(pastKeySizes[3] == headSize);
}
if (hasPastValue)
{
auto pastValueSizes = m_inputTensorDescs[dmlPastValueIndex].GetSizes();
ML_CHECK_VALID_ARGUMENT(pastValueSizes[0] == batchSize);
ML_CHECK_VALID_ARGUMENT(pastValueSizes[1] == numHeads);
ML_CHECK_VALID_ARGUMENT(pastValueSizes[2] == pastSequenceLength);
ML_CHECK_VALID_ARGUMENT(pastValueSizes[3] == headSize);
}
std::vector<DML_TENSOR_DESC> inputDescs = GetDmlInputDescs();
std::vector<DML_TENSOR_DESC> outputDescs = GetDmlOutputDescs();
DML_MULTIHEAD_ATTENTION_OPERATOR_DESC mhaDesc = {};
mhaDesc.QueryTensor = stackedQkv ? nullptr : &inputDescs[dmlQueryIndex];
mhaDesc.KeyTensor = hasKey ? &inputDescs[dmlKeyIndex] : nullptr;
mhaDesc.ValueTensor = hasValue ? &inputDescs[dmlValueIndex] : nullptr;
mhaDesc.StackedKeyValueTensor = stackedKv ? &inputDescs[dmlStackedKeyValueIndex] : nullptr;
mhaDesc.StackedQueryKeyValueTensor = stackedQkv ? &inputDescs[dmlStackedQueryKeyValueIndex] : nullptr;
mhaDesc.BiasTensor = hasBias ? &inputDescs[dmlBiasIndex] : nullptr;
mhaDesc.MaskTensor = hasMask ? &inputDescs[dmlMaskIndex] : nullptr;
mhaDesc.RelativePositionBiasTensor = hasRelativePositionBias ? &inputDescs[dmlRelativePositionBiasIndex] : nullptr;
mhaDesc.PastKeyTensor = hasPastKey ? &inputDescs[dmlPastKeyIndex] : nullptr;
mhaDesc.PastValueTensor = hasPastValue ? &inputDescs[dmlPastValueIndex] : nullptr;
mhaDesc.OutputTensor = &outputDescs[outputIndex];
mhaDesc.OutputPresentKeyTensor = hasPresentKeyOutput ? &outputDescs[outputPresentKeyIndex] : nullptr;
mhaDesc.OutputPresentValueTensor = hasPresentValueOutput ? &outputDescs[outputPresentValueIndex] : nullptr;
mhaDesc.Scale = kernelCreationContext.GetOptionalAttribute<float>(AttrName::Scale, gsl::narrow_cast<float>(1.0f / std::sqrt(headSize)));
mhaDesc.MaskFilterValue = kernelCreationContext.GetOptionalAttribute<float>(AttrName::MaskFilterValue, -10'000.0f);
mhaDesc.HeadCount = numHeads;
mhaDesc.MaskType = maskType;
DML_OPERATOR_DESC opDesc = { DML_OPERATOR_MULTIHEAD_ATTENTION, &mhaDesc };
SetDmlOperatorDesc(opDesc, kernelCreationContext);
}
};
DML_OP_DEFINE_CREATION_FUNCTION(MultiHeadAttention, DmlOperatorMultiHeadAttention);
} // namespace Dml

View file

@ -386,6 +386,7 @@ DML_OP_EXTERN_CREATION_FUNCTION(Trilu);
DML_OP_EXTERN_CREATION_FUNCTION(Shape);
DML_OP_EXTERN_CREATION_FUNCTION(Size);
DML_OP_EXTERN_CREATION_FUNCTION(Attention);
DML_OP_EXTERN_CREATION_FUNCTION(MultiHeadAttention);
DML_OP_EXTERN_CREATION_FUNCTION(NonZero);
DML_OP_EXTERN_CREATION_FUNCTION(QuickGelu);
@ -871,8 +872,10 @@ constexpr static OperatorRegistrationInformation operatorRegistrationInformation
{REG_INFO_MS( 1, Gelu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MS( 1, BiasGelu, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MS( 1, FusedMatMul, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MS( 1, FusedMatMulActivation, typeNameListDefault, supportedTypeListFloat16to32, DmlGraphSupport::Supported)},
{REG_INFO_MS( 1, QLinearSigmoid, typeNameListDefault, supportedTypeListQLinearSigmoid, DmlGraphSupport::Supported, requiredConstantCpuInputs(), std::nullopt, QueryQLinearSigmoid)},
{REG_INFO_MS( 1, Attention, typeNameListAttention, supportedTypeListAttention, DmlGraphSupport::Supported, requiredConstantCpuInputs(), std::nullopt, QueryAttention)},
{REG_INFO_MS( 1, MultiHeadAttention, typeNameListAttention, supportedTypeListAttention, DmlGraphSupport::Supported)},
{REG_INFO( 10, IsInf, typeNameListTwo, supportedTypeListIsInf, DmlGraphSupport::Supported)},
{REG_INFO( 10, Mod, typeNameListDefault, supportedTypeListNumericDefault, DmlGraphSupport::Supported)},

View file

@ -364,19 +364,10 @@ namespace Dml
int onnxAxis = HandleNegativeAxis(kernelInfo.GetOptionalAttribute<int>(AttrName::GraphFusedAxis, -1), onnxDimCount);
auto dmlAdjustedAxis = GetDmlAdjustedAxis(onnxAxis, onnxDimCount, kernelInfo.GetTensorShapeDescription().GetInputTensorDimensionCount(0));
// If the axis is supported by Softmax, use this version instead since it's more likely to be supported by metacommands
if (dmlAdjustedAxis == onnxDimCount - 1)
{
activation.desc.activationType = DML_OPERATOR_ACTIVATION_SOFTMAX;
}
else
{
activation.desc.activationType = DML_OPERATOR_ACTIVATION_SOFTMAX1;
activation.dmlAxes.push_back(dmlAdjustedAxis);
activation.desc.params.softmax1.Axes = activation.dmlAxes.data();
activation.desc.params.softmax1.AxisCount = gsl::narrow_cast<uint32_t>(activation.dmlAxes.size());
}
activation.desc.activationType = DML_OPERATOR_ACTIVATION_SOFTMAX1;
activation.dmlAxes.push_back(dmlAdjustedAxis);
activation.desc.params.softmax1.Axes = activation.dmlAxes.data();
activation.desc.params.softmax1.AxisCount = gsl::narrow_cast<uint32_t>(activation.dmlAxes.size());
}
else
{

View file

@ -113,6 +113,8 @@ namespace AttrName
static constexpr const char* FusedBeta = "fused_beta";
static constexpr const char* FusedGamma = "fused_gamma";
static constexpr const char* FusedRatio = "fused_ratio";
static constexpr const char* MaskFilterValue = "mask_filter_value";
static constexpr const char* DoRotary = "do_rotary";
static constexpr const char* Activation = "activation";
static constexpr const char* Groups = "groups";

View file

@ -2536,6 +2536,114 @@ namespace OperatorHelper
return outputShapes;
}
std::vector<EdgeShapes> MultiHeadAttentionHelper::GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const
{
ML_CHECK_VALID_ARGUMENT(shapeInfo.GetInputCount() >= 1);
auto queryShape = shapeInfo.GetInputTensorShape(0);
ML_CHECK_VALID_ARGUMENT(queryShape.size() == 3 || queryShape.size() == 5);
const uint32_t batchSize = queryShape[0];
const uint32_t sequenceLength = queryShape[1];
uint32_t kvSequenceLength = 0;
uint32_t vHiddenSize = 0;
uint32_t headSize = 0;
if (shapeInfo.IsInputValid(2))
{
auto valueShape = shapeInfo.GetInputTensorShape(2);
ML_CHECK_VALID_ARGUMENT(queryShape.size() == 3);
headSize = queryShape[2] / m_numHeads;
if (valueShape.size() == 3)
{
kvSequenceLength = valueShape[1];
vHiddenSize = valueShape[2];
}
else
{
ML_CHECK_VALID_ARGUMENT(valueShape.size() == 4);
const uint32_t vHeadSize = valueShape[3];
kvSequenceLength = valueShape[2];
vHiddenSize = vHeadSize * m_numHeads;
}
}
else if (shapeInfo.IsInputValid(1))
{
auto keyShape = shapeInfo.GetInputTensorShape(1);
ML_CHECK_VALID_ARGUMENT(keyShape.size() == 5);
kvSequenceLength = keyShape[1];
vHiddenSize = queryShape[2];
headSize = keyShape[4];
}
else
{
ML_CHECK_VALID_ARGUMENT(queryShape.size() == 5);
kvSequenceLength = queryShape[1];
headSize = queryShape[4];
vHiddenSize = headSize * m_numHeads;
}
std::vector<EdgeShapes> outputShapes(3);
outputShapes[0] = EdgeShapes({batchSize, sequenceLength, vHiddenSize});
uint32_t totalSequenceLength = kvSequenceLength;
if (shapeInfo.IsInputValid(6))
{
ML_CHECK_VALID_ARGUMENT(shapeInfo.GetInputTensorDimensionCount(6) == 4);
const uint32_t pastSequenceLength = shapeInfo.GetInputTensorShape(6)[2];
totalSequenceLength += pastSequenceLength;
}
if (shapeInfo.IsOutputValid(1))
{
outputShapes[1] = EdgeShapes({batchSize, m_numHeads, totalSequenceLength, headSize});
}
if (shapeInfo.IsOutputValid(2))
{
outputShapes[2] = EdgeShapes({batchSize, m_numHeads, totalSequenceLength, headSize});
}
return outputShapes;
}
void MultiHeadAttentionHelper::Initialize(const IKernelInformationAdapter& kernelInformation)
{
m_numHeads = gsl::narrow_cast<uint32_t>(kernelInformation.GetAttributes().GetAttribute<int64_t>(AttrName::NumHeads));
}
std::vector<EdgeShapes> AttentionHelper::GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const
{
ML_CHECK_VALID_ARGUMENT(shapeInfo.GetInputCount() >= 2);
auto queryShape = shapeInfo.GetInputTensorShape(0);
ML_CHECK_VALID_ARGUMENT(queryShape.size() == 3);
auto weightShape = shapeInfo.GetInputTensorShape(1);
ML_CHECK_VALID_ARGUMENT(weightShape.size() == 2);
if (m_qkvHiddenSizes.empty())
{
ML_CHECK_VALID_ARGUMENT(weightShape[1] % 3 == 0);
}
else
{
ML_CHECK_VALID_ARGUMENT(m_qkvHiddenSizes.size() == 3);
}
const uint32_t batchSize = queryShape[0];
const uint32_t sequenceLength = queryShape[1];
const uint32_t vHiddenSize = m_qkvHiddenSizes.empty() ? weightShape[1] / 3 : m_qkvHiddenSizes[2];
return { EdgeShapes({batchSize, sequenceLength, vHiddenSize}) };
}
void AttentionHelper::Initialize(const IKernelInformationAdapter& kernelInformation)
{
m_qkvHiddenSizes = kernelInformation.GetAttributes().GetOptionalAttributeVectorInt32(AttrName::QkvHiddenSizes);
}
std::vector<EdgeShapes> SkipLayerNormHelper::GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const
{
ML_CHECK_VALID_ARGUMENT(shapeInfo.GetInputCount() >= 3);

View file

@ -1419,6 +1419,38 @@ public:
std::vector<EdgeShapes> GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const;
};
class MultiHeadAttentionHelper
{
public:
template <typename Info_t, typename Shape_t>
MultiHeadAttentionHelper(const Info_t& info, const Shape_t& shapeInfo)
{
Initialize(KernelInformationAdapter(info));
}
std::vector<EdgeShapes> GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const;
private:
void Initialize(const IKernelInformationAdapter& kernelInformation);
uint32_t m_numHeads;
};
class AttentionHelper
{
public:
template <typename Info_t, typename Shape_t>
AttentionHelper(const Info_t& info, const Shape_t& shapeInfo)
{
Initialize(KernelInformationAdapter(info));
}
std::vector<EdgeShapes> GetOutputShapes(const MLShapeInferenceContext& shapeInfo) const;
private:
void Initialize(const IKernelInformationAdapter& kernelInformation);
std::vector<int32_t> m_qkvHiddenSizes;
};
class SkipLayerNormHelper
{
public:
@ -1554,7 +1586,8 @@ using ShapeInferenceHelper_Affine = GetOutputShapeAsInputShapeHelper;
using ShapeInferenceHelper_QuantizeLinear = GetOutputShapeAsInputShapeHelper;
using ShapeInferenceHelper_DequantizeLinear = GetOutputShapeAsInputShapeHelper;
using ShapeInferenceHelper_QLinearSigmoid = GetOutputShapeAsInputShapeHelper;
using ShapeInferenceHelper_Attention = GetOutputShapeAsInputShapeHelper;
using ShapeInferenceHelper_Attention = AttentionHelper;
using ShapeInferenceHelper_MultiHeadAttention = MultiHeadAttentionHelper;
using ShapeInferenceHelper_Sign = GetBroadcastedOutputShapeHelper;
using ShapeInferenceHelper_IsNaN = GetBroadcastedOutputShapeHelper;
using ShapeInferenceHelper_Erf = GetBroadcastedOutputShapeHelper;

View file

@ -423,6 +423,7 @@ namespace OperatorHelper
static const int sc_sinceVer_FusedMatMulActivation = 1;
static const int sc_sinceVer_QLinearSigmoid = 1;
static const int sc_sinceVer_Attention = 1;
static const int sc_sinceVer_MultiHeadAttention = 1;
static const int sc_sinceVer_SkipLayerNormalization = 1;
static const int sc_sinceVer_EmbedLayerNormalization = 1;
static const int sc_sinceVer_BiasSplitGelu = 1;

View file

@ -58,6 +58,7 @@ static void RunAttentionTest(
const bool disable_cpu = false,
const bool disable_cuda = false,
const bool disable_rocm = false,
const bool disable_dml = false,
std::vector<int32_t> qkv_sizes = {},
const std::vector<float>& relative_position_bias_data = {},
int kv_sequence_length = 0,
@ -72,9 +73,10 @@ static void RunAttentionTest(
bool enable_cuda = HasCudaEnvironment(min_cuda_architecture) && !is_weights_constant && !disable_cuda;
bool enable_rocm = (nullptr != DefaultRocmExecutionProvider().get()) && !is_weights_constant && !disable_rocm;
bool enable_cpu = (nullptr != DefaultCpuExecutionProvider().get()) && !use_float16 && !disable_cpu;
bool enable_dml = (nullptr != DefaultDmlExecutionProvider().get()) && !disable_dml;
int head_size = hidden_size / number_of_heads;
if (enable_cpu || enable_cuda || enable_rocm) {
if (enable_cpu || enable_cuda || enable_rocm || enable_dml) {
OpTester tester("Attention", 1, onnxruntime::kMSDomain);
tester.AddAttribute<int64_t>("num_heads", static_cast<int64_t>(number_of_heads));
tester.AddAttribute<int64_t>("unidirectional", static_cast<int64_t>(is_unidirectional ? 1 : 0));
@ -242,6 +244,12 @@ static void RunAttentionTest(
execution_providers.push_back(DefaultCpuExecutionProvider());
tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}
if (enable_dml) {
std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultDmlExecutionProvider());
tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}
}
}
@ -267,6 +275,7 @@ static void RunAttentionTest(
const bool disable_cpu = false,
const bool disable_cuda = false,
const bool disable_rocm = false,
const bool disable_dml = false,
const std::vector<int32_t> qkv_sizes = {},
const std::vector<float>& relative_position_bias_data = {},
int kv_sequence_length = 0,
@ -277,13 +286,13 @@ static void RunAttentionTest(
batch_size, sequence_length, hidden_size, number_of_heads,
use_float16, is_unidirectional, use_past_state, past_sequence_length,
past_data, present_data, mask_type, input_hidden_size, max_sequence_length,
disable_cpu, disable_cuda, disable_rocm, qkv_sizes, relative_position_bias_data,
disable_cpu, disable_cuda, disable_rocm, disable_dml, qkv_sizes, relative_position_bias_data,
kv_sequence_length, past_present_share_buffer, use_scale, do_neox_rotary);
RunAttentionTest(input_data, weights_data, true, bias_data, mask_index_data, output_data,
batch_size, sequence_length, hidden_size, number_of_heads,
use_float16, is_unidirectional, use_past_state, past_sequence_length,
past_data, present_data, mask_type, input_hidden_size, max_sequence_length,
disable_cpu, disable_cuda, disable_rocm, qkv_sizes, relative_position_bias_data,
disable_cpu, disable_cuda, disable_rocm, disable_dml, qkv_sizes, relative_position_bias_data,
kv_sequence_length, past_present_share_buffer, use_scale, do_neox_rotary);
}
@ -354,7 +363,7 @@ TEST(AttentionTest, AttentionBatch1WithQKVAttr1) {
RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
batch_size, sequence_length, hidden_size, number_of_heads,
false, false, false, 0, nullptr, nullptr, AttentionMaskType::MASK_1D_KEY_SEQ_LEN, 0,
0, false, false, disable_rocm, qkv_sizes);
0, false, false, disable_rocm, false, qkv_sizes);
}
TEST(AttentionTest, AttentionBatch1WithQKVAttr2) {
@ -392,7 +401,7 @@ TEST(AttentionTest, AttentionBatch1WithQKVAttr2) {
RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
batch_size, sequence_length, hidden_size, number_of_heads,
false, false, false, 0, nullptr, nullptr, AttentionMaskType::MASK_1D_KEY_SEQ_LEN, 0,
0, false, false, disable_rocm, qkv_sizes);
0, false, false, disable_rocm, false, qkv_sizes);
}
TEST(AttentionTest, AttentionBatch1RelativePositionBias) {
@ -429,10 +438,11 @@ TEST(AttentionTest, AttentionBatch1RelativePositionBias) {
constexpr bool disable_cpu = false;
constexpr bool disable_cuda = false;
constexpr bool disable_rocm = false;
constexpr bool disable_dml = false;
RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
batch_size, sequence_length, hidden_size, number_of_heads,
false, false, false, 0, nullptr, nullptr, AttentionMaskType::MASK_1D_KEY_SEQ_LEN, 0,
0, disable_cpu, disable_cuda, disable_rocm, qkv_sizes, relative_position_bias);
0, disable_cpu, disable_cuda, disable_rocm, disable_dml, qkv_sizes, relative_position_bias);
}
TEST(AttentionTest, AttentionBatch2RelativePositionBias) {
@ -474,10 +484,11 @@ TEST(AttentionTest, AttentionBatch2RelativePositionBias) {
constexpr bool disable_cpu = false;
constexpr bool disable_cuda = false;
constexpr bool disable_rocm = false;
constexpr bool disable_dml = false;
RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
batch_size, sequence_length, hidden_size, number_of_heads,
false, false, false, 0, nullptr, nullptr, AttentionMaskType::MASK_1D_KEY_SEQ_LEN, 0,
0, disable_cpu, disable_cuda, disable_rocm, qkv_sizes, relative_position_bias);
0, disable_cpu, disable_cuda, disable_rocm, disable_dml, qkv_sizes, relative_position_bias);
}
TEST(AttentionTest, AttentionBatch1_Float16) {
@ -817,10 +828,14 @@ void RawAttentionEmptyPastState(bool past_present_share_buffer) {
batch_size, sequence_length, hidden_size, number_of_heads, false, is_unidirectional,
use_past_state, past_sequence_length, &past_data, &present_data);
} else {
// TODO: Unskip when fixed #41968513
// DML doesn't support past_present_share_buffer for Attention yet
constexpr bool disable_dml = true;
RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
batch_size, sequence_length, hidden_size, number_of_heads, false, is_unidirectional,
use_past_state, past_sequence_length, &past_data, &present_data,
AttentionMaskType::MASK_1D_KEY_SEQ_LEN, 0, sequence_length, true, false, true, {}, {}, 0,
AttentionMaskType::MASK_1D_KEY_SEQ_LEN, 0, sequence_length, true, false, true, disable_dml, {}, {}, 0,
true);
}
}
@ -1075,11 +1090,15 @@ void RawAttentionPastStateBatch1(bool past_present_share_buffer) {
batch_size, sequence_length, hidden_size, number_of_heads, false, is_unidirectional,
use_past_state, past_sequence_length, &past_data, &present_data);
} else {
// TODO: Unskip when fixed #41968513
// DML doesn't support past_present_share_buffer for Attention yet
constexpr bool disable_dml = true;
RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
batch_size, sequence_length, hidden_size, number_of_heads, false, is_unidirectional,
use_past_state, past_sequence_length, &past_data, &present_data,
AttentionMaskType::MASK_1D_KEY_SEQ_LEN, 0, past_sequence_length + sequence_length + 4,
true, false, true, {}, {}, 0, true);
true, false, true, disable_dml, {}, {}, 0, true);
}
}
@ -1204,11 +1223,15 @@ void RawAttentionPastStateBatch2(bool past_present_share_buffer) {
batch_size, sequence_length, hidden_size, number_of_heads, false, is_unidirectional,
use_past_state, past_sequence_length, &past_data, &present_data);
} else {
// TODO: Unskip when fixed #41968513
// DML doesn't support past_present_share_buffer for Attention yet
constexpr bool disable_dml = true;
RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
batch_size, sequence_length, hidden_size, number_of_heads, false, is_unidirectional,
use_past_state, past_sequence_length, &past_data, &present_data,
AttentionMaskType::MASK_1D_KEY_SEQ_LEN, 0, past_sequence_length + sequence_length,
true, false, true, {}, {}, 0, true);
true, false, true, disable_dml, {}, {}, 0, true);
}
}
@ -1324,12 +1347,16 @@ void RawAttentionPastStateBatch2WithPadding(bool past_present_share_buffer) {
use_past_state, past_sequence_length, &past_data, &present_data,
AttentionMaskType::MASK_1D_END_START);
} else {
// TODO: Unskip when fixed #41968513
// DML doesn't support past_present_share_buffer for Attention yet
constexpr bool disable_dml = true;
RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
batch_size, sequence_length, hidden_size, number_of_heads, false, is_unidirectional,
use_past_state, past_sequence_length, &past_data, &present_data,
AttentionMaskType::MASK_1D_END_START,
0, past_sequence_length + sequence_length + 4,
true, false, true, {}, {}, 0, true);
true, false, true, disable_dml, {}, {}, 0, true);
}
}
@ -1716,7 +1743,7 @@ TEST(AttentionTest, AttentionWithNormFactor) {
batch_size, sequence_length, hidden_size, number_of_heads,
use_float16, is_unidirectional, use_past_state, past_sequence_length, past_data, present_data,
AttentionMaskType::MASK_2D_KEY_PADDING, 0 /*input_hidden_size*/, 0 /*max_sequence_length*/,
false /*disable_cpu*/, false /*disable_cuda*/, true /*disable_rocm*/, {} /*qkv_sizes*/,
false /*disable_cpu*/, false /*disable_cuda*/, true /*disable_rocm*/, false /*disable_dml*/, {} /*qkv_sizes*/,
{} /*relative_position_bias_data*/, 0 /*kv_sequence_length*/, false /*past_present_share_buffer*/,
true /*use_scale*/);
}
@ -1757,11 +1784,16 @@ TEST(AttentionTest, AttentionWithNeoXRotaryEmbedding) {
int past_sequence_length = 0;
const std::vector<float>* past_data = nullptr;
const std::vector<float>* present_data = nullptr;
// TODO: Unskip when fixed #41968513
// DML doesn't support do_rotary for Attention yet
constexpr bool disable_dml = true;
RunAttentionTest(input_data, weight_data, bias_data, mask_index_data, output_data,
batch_size, sequence_length, hidden_size, number_of_heads,
use_float16, is_unidirectional, use_past_state, past_sequence_length, past_data, present_data,
AttentionMaskType::MASK_2D_KEY_PADDING, 0 /*input_hidden_size*/, 0 /*max_sequence_length*/,
true /*disable_cpu*/, false /*disable_cuda*/, true /*disable_rocm*/, {} /*qkv_sizes*/,
true /*disable_cpu*/, false /*disable_cuda*/, true /*disable_rocm*/, disable_dml, {} /*qkv_sizes*/,
{} /*relative_position_bias_data*/, 0 /*kv_sequence_length*/, false /*past_present_share_buffer*/,
true /*use_scale*/, true /*use_neox_rotary_embedding*/);
}
@ -2191,7 +2223,8 @@ static void RunModelWithRandomInput(
bool enable_cuda = HasCudaEnvironment(is_float16 ? 530 : 0);
bool enable_rocm = (nullptr != DefaultRocmExecutionProvider().get());
bool enable_cpu = (nullptr != DefaultCpuExecutionProvider().get() && !is_float16);
if (enable_cuda || enable_rocm) {
bool enable_dml = (nullptr != DefaultDmlExecutionProvider().get());
if (enable_cuda || enable_rocm || enable_dml) {
OpTester test("Attention", 1, onnxruntime::kMSDomain);
test.AddAttribute<int64_t>("num_heads", num_heads);
if (is_float16) {
@ -2208,6 +2241,8 @@ static void RunModelWithRandomInput(
std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
if (enable_cuda) {
execution_providers.push_back(DefaultCudaExecutionProvider());
} else if (enable_dml) {
execution_providers.push_back(DefaultDmlExecutionProvider());
} else {
execution_providers.push_back(DefaultRocmExecutionProvider());
}

View file

@ -7,7 +7,7 @@
namespace onnxruntime {
namespace test {
#ifndef _MSC_VER
#if !defined(_MSC_VER) || defined(USE_DML)
void GetWeight_64_3_64(std::vector<float>& weight_data) {
weight_data = {
-0.004707f, -0.006775f, 0.0009236f, 0.003067f, -0.00806f, 0.00779f, 0.0004425f, 0.00846f, 0.00048f,

View file

@ -43,7 +43,7 @@ struct AttentionTestData {
};
// Disable some tests in Windows since prefast build might crash with large test data.
#ifndef _MSC_VER
#if !defined(_MSC_VER) || defined(USE_DML)
// Return packed weights and bias for input projection.
void GetAttentionWeight(std::vector<float>& weight_data, int elements = 64 * 3 * 64, int offset = 0, int step = 1);
void GetAttentionBias(std::vector<float>& bias_data, int elements = 3 * 64, int offset = 0, int step = 1);

View file

@ -55,8 +55,8 @@ static void RunMultiHeadAttentionTest(
bool use_float16 = false,
bool disable_cpu = false, // some cases not supported in cpu right now.
bool disable_cuda = false,
bool disable_rocm = DISABLE_ROCM) // not supported in rocm right now.
{
bool disable_rocm = DISABLE_ROCM, // not supported in rocm right now.
bool disable_dml = false) {
kv_sequence_length = (kv_sequence_length == 0 ? sequence_length : kv_sequence_length);
int min_cuda_architecture = use_float16 ? 750 : 0;
@ -64,8 +64,9 @@ static void RunMultiHeadAttentionTest(
// rocm mha is required to work with TunableOp Enabled
bool enable_rocm = (nullptr != DefaultRocmExecutionProvider(/*test_tunable_op=*/true).get()) && !disable_rocm;
bool enable_cpu = (nullptr != DefaultCpuExecutionProvider().get()) && !use_float16 && !disable_cpu;
bool enable_dml = (nullptr != DefaultDmlExecutionProvider().get()) && !disable_dml;
if (enable_cpu || enable_cuda || enable_rocm) {
if (enable_cpu || enable_cuda || enable_rocm || enable_dml) {
OpTester tester("MultiHeadAttention", 1, onnxruntime::kMSDomain);
tester.AddAttribute<int64_t>("num_heads", static_cast<int64_t>(num_heads));
tester.AddAttribute<float>("mask_filter_value", static_cast<float>(-10000.0f));
@ -255,6 +256,12 @@ static void RunMultiHeadAttentionTest(
execution_providers.push_back(DefaultCpuExecutionProvider());
tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}
if (enable_dml) {
std::vector<std::unique_ptr<IExecutionProvider>> execution_providers;
execution_providers.push_back(DefaultDmlExecutionProvider());
tester.Run(OpTester::ExpectResult::kExpectSuccess, "", {}, nullptr, &execution_providers);
}
}
}
@ -284,7 +291,8 @@ static void RunMultiHeadAttentionKernel(
bool is_static_kv = true,
bool disable_cpu = false, // some cases not supported in cpu right now.
bool disable_cuda = false,
bool disable_rocm = DISABLE_ROCM) {
bool disable_rocm = DISABLE_ROCM,
bool disable_dml = false) {
if (kernel_type == AttentionKernelType::AttentionKernel_Default) {
ScopedEnvironmentVariables scoped_env_vars{
EnvVarMap{
@ -296,7 +304,7 @@ static void RunMultiHeadAttentionKernel(
query_data, key_data, value_data, kv_data, qkv_data, bias_data, rel_pos_bias_data,
past_key_data, past_value_data, present_key_data, present_value_data, key_padding_mask_data,
mask_type, output_data, num_heads, batch_size, sequence_length, kv_sequence_length,
hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm);
hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm, disable_dml);
return;
}
@ -311,7 +319,7 @@ static void RunMultiHeadAttentionKernel(
query_data, key_data, value_data, kv_data, qkv_data, bias_data, rel_pos_bias_data,
past_key_data, past_value_data, present_key_data, present_value_data, key_padding_mask_data,
mask_type, output_data, num_heads, batch_size, sequence_length, kv_sequence_length,
hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm);
hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm, disable_dml);
return;
}
@ -326,7 +334,7 @@ static void RunMultiHeadAttentionKernel(
query_data, key_data, value_data, kv_data, qkv_data, bias_data, rel_pos_bias_data,
past_key_data, past_value_data, present_key_data, present_value_data, key_padding_mask_data,
mask_type, output_data, num_heads, batch_size, sequence_length, kv_sequence_length,
hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm);
hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm, disable_dml);
return;
}
@ -342,7 +350,7 @@ static void RunMultiHeadAttentionKernel(
query_data, key_data, value_data, kv_data, qkv_data, bias_data, rel_pos_bias_data,
past_key_data, past_value_data, present_key_data, present_value_data, key_padding_mask_data,
mask_type, output_data, num_heads, batch_size, sequence_length, kv_sequence_length,
hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm);
hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm, disable_dml);
return;
}
#endif
@ -358,7 +366,7 @@ static void RunMultiHeadAttentionKernel(
query_data, key_data, value_data, kv_data, qkv_data, bias_data, rel_pos_bias_data,
past_key_data, past_value_data, present_key_data, present_value_data, key_padding_mask_data,
mask_type, output_data, num_heads, batch_size, sequence_length, kv_sequence_length,
hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm);
hidden_size, v_hidden_size, is_static_kv, use_float16, disable_cpu, disable_cuda, disable_rocm, disable_dml);
}
}
@ -444,7 +452,7 @@ static void RunMultiHeadAttentionTests(AttentionTestData& data, bool disable_cpu
}
}
#ifndef _MSC_VER
#if !defined(_MSC_VER) || defined(USE_DML)
// Test fused cross attention kernel
// It requires head_size > 32 and head_size <= 64 for T4 GPU; hidden_size == v_hidden_size.
TEST(MultiHeadAttentionTest, CrossAttention_Batch2_HeadSize40) {

View file

@ -342,7 +342,7 @@ struct TensorCheck<MLFloat16> {
#if defined(USE_TENSORRT) || defined(ENABLE_TRAINING_CORE) || defined(USE_CUDA) || defined(USE_ROCM)
threshold = 0.005f;
#elif defined(USE_DML)
threshold = 0.008f;
threshold = 0.02f;
#endif
for (int i = 0; i < size; ++i) {
if (std::isnan(f_expected[i])) {

View file

@ -1,6 +1,6 @@
<?xml version="1.0" encoding="utf-8"?>
<packages>
<package id="Microsoft.AI.DirectML" version="1.11.0" targetFramework="native" />
<package id="Microsoft.AI.DirectML" version="1.12.0" targetFramework="native" />
<package id="Microsoft.Windows.CppWinRT" version="2.0.201201.7" targetFramework="native" />
<package id="google.protobuf.tools" version="3.21.12" targetFramework="native" />
</packages>

View file

@ -677,6 +677,8 @@ def parse_arguments():
parser.add_argument("--use_cache", action="store_true", help="Use compiler cache in CI")
parser.add_argument("--use_lock_free_queue", action="store_true", help="Use lock-free task queue for threadpool.")
if not is_windows():
parser.add_argument(
"--allow_running_as_root",
@ -1307,6 +1309,9 @@ def generate_build_tree(
if args.use_azure:
add_default_definition(cmake_extra_defines, "onnxruntime_USE_AZURE", "ON")
if args.use_lock_free_queue:
add_default_definition(cmake_extra_defines, "onnxruntime_USE_LOCK_FREE_QUEUE", "ON")
cmake_args += [f"-D{define}" for define in cmake_extra_defines]
cmake_args += cmake_extra_args

View file

@ -200,7 +200,7 @@ stages:
BuildConfig: 'RelWithDebInfo'
EnvSetupScript: setup_env_azure.bat
buildArch: x64
additionalBuildFlags: --use_azure
additionalBuildFlags: --use_azure --use_lock_free_queue
msbuildPlatform: x64
isX86: false
job_name_suffix: x64_release_azure

View file

@ -190,7 +190,7 @@ def generate_repo_url(line_list, repo_url, commit_id):
def generate_dependencies(xml_text, package_name, version):
dml_dependency = '<dependency id="Microsoft.AI.DirectML" version="1.11.0"/>'
dml_dependency = '<dependency id="Microsoft.AI.DirectML" version="1.12.0"/>'
if package_name == "Microsoft.AI.MachineLearning":
xml_text.append("<dependencies>")