Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
67 commits
Select commit Hold shift + click to select a range
f314287
Add support for creating optimized whisper ONNX models without beam s…
kunal-vaishnavi Apr 26, 2024
6a44f72
Fix incorrect dynamic axes labels
kunal-vaishnavi Apr 26, 2024
58ec5eb
Fix fusion breaks for OpenAI implementation of Whisper
kunal-vaishnavi May 3, 2024
4c228ea
Merge branch 'main' into kvaishnavi/whisper-separate-export
kunal-vaishnavi Jun 12, 2024
dd20876
Merge branch 'main' into kvaishnavi/whisper-separate-export
kunal-vaishnavi Jul 23, 2024
b13cb22
Comment out DMMHA case temporarily
kunal-vaishnavi Jul 23, 2024
31db1a0
Replace MHA with DMMHA
kunal-vaishnavi Jul 29, 2024
3b92432
Merge branch 'main' into kvaishnavi/whisper-separate-export
kunal-vaishnavi Aug 26, 2024
7bb79f3
Debugging beam search output
kunal-vaishnavi Sep 6, 2024
14b7e77
Initial commit for new export
kunal-vaishnavi Oct 22, 2024
fa345fe
Add parity check after export and optimization
kunal-vaishnavi Oct 22, 2024
e050dea
Fix multiple attention kernel invocations
kunal-vaishnavi Nov 2, 2024
bf87062
Make output Q*K values optional
kunal-vaishnavi Nov 4, 2024
17fa0ab
Fix batch size check for cache indirection
kunal-vaishnavi Nov 6, 2024
52aeb58
Save checkpoint for working solution
kunal-vaishnavi Nov 15, 2024
240fe3b
Clean up code
kunal-vaishnavi Nov 17, 2024
ae98085
Fix string dumping
kunal-vaishnavi Nov 20, 2024
3d2c8fe
Fix out_qk dtype issue for half input case.
mindest Nov 20, 2024
287151f
Remove type cast for output QK
kunal-vaishnavi Nov 21, 2024
0805d1d
Enable release mode build
kunal-vaishnavi Dec 4, 2024
b629903
Make QK output dtype independent of attention dtype
kunal-vaishnavi Dec 9, 2024
648b389
Add batched jump times export
kunal-vaishnavi Dec 9, 2024
a6c6ee8
Get batched jump times ONNX model with parity check
kunal-vaishnavi Dec 12, 2024
c0a6ce4
Save checkpoint for working solution
kunal-vaishnavi Dec 21, 2024
008eeb9
Merge branch 'main' into kvaishnavi/whisper-separate-export
kunal-vaishnavi Dec 22, 2024
158d0a8
Fix build after merge
kunal-vaishnavi Dec 22, 2024
02cb5be
Fix model with beam search op
kunal-vaishnavi Dec 23, 2024
2acd593
Get model impl and beam search op export combinations working
kunal-vaishnavi Dec 25, 2024
612eb0c
Enable separate export of encoder and decoder init
kunal-vaishnavi Dec 25, 2024
f2d78fd
Add tests for multiple export types to CIs
kunal-vaishnavi Dec 25, 2024
cb93517
Update folder and file names in Whisper README
kunal-vaishnavi Dec 25, 2024
6da11ec
Add FP32 CPU DMMHA support
kunal-vaishnavi Dec 28, 2024
9640736
Add unit tests
kunal-vaishnavi Jan 8, 2025
75a342a
Merge branch 'main' into kvaishnavi/whisper-separate-export
kunal-vaishnavi Jan 24, 2025
7fe6b05
Change debug message for PrepareQkv
kunal-vaishnavi Jan 25, 2025
8620168
Fix seqlens_k after merge
kunal-vaishnavi Jan 29, 2025
b0a732b
Merge branch 'main' into kvaishnavi/whisper-separate-export
kunal-vaishnavi Jan 29, 2025
23808f7
Add changes suggested by linter
kunal-vaishnavi Jan 31, 2025
906023d
Fix bug in FP32 CPU jump times model
kunal-vaishnavi Mar 6, 2025
3ed4bf2
Merge branch 'main' into kvaishnavi/whisper-separate-export
kunal-vaishnavi Mar 6, 2025
fae3dd8
Add changes from PR feedback
kunal-vaishnavi Mar 6, 2025
3c84fd4
Merge branch 'main' into kvaishnavi/whisper-separate-export
kunal-vaishnavi Mar 10, 2025
f3003fb
Compare token ids outputs of various shapes
kunal-vaishnavi Mar 10, 2025
f8c04fe
Fix MHA unit test failures
kunal-vaishnavi Mar 11, 2025
52e0fd0
Fix Whisper fusion tests
kunal-vaishnavi Mar 12, 2025
78a0787
Remove debugging code line
kunal-vaishnavi Mar 12, 2025
4f68e40
Fix more CI unit tests
kunal-vaishnavi Mar 12, 2025
33183a7
Fix CI build errors
kunal-vaishnavi Mar 13, 2025
bd38ccc
Fix more CI build errors
kunal-vaishnavi Mar 13, 2025
65b1739
Add ninja to docker image and update docs
kunal-vaishnavi Mar 13, 2025
53a470c
Fix typo with package name
kunal-vaishnavi Mar 13, 2025
130626f
Upgrade to CUDA 12.1 in CIs
kunal-vaishnavi Mar 13, 2025
9e20aea
Attempt to upgrade to CUDA 12.4
kunal-vaishnavi Mar 13, 2025
f6eabd4
Revert back to CUDA 11.8 in CIs
kunal-vaishnavi Mar 14, 2025
e443d70
Fix typo in TRT version when reverting
kunal-vaishnavi Mar 14, 2025
ab96683
Merge branch 'main' into kvaishnavi/whisper-separate-export
kunal-vaishnavi Mar 14, 2025
11a69fc
Add changes based on PR feedback
kunal-vaishnavi Mar 14, 2025
f04bd0b
Merge branch 'main' into kvaishnavi/whisper-separate-export
kunal-vaishnavi Mar 14, 2025
0adafe7
Rename from FT causal attention to decoder attention
kunal-vaishnavi Mar 14, 2025
460e7e0
Fix Python linter error
kunal-vaishnavi Mar 14, 2025
3ed3a47
Update buffer sharing definition
kunal-vaishnavi Mar 14, 2025
09d9fef
Update MHA op spec
kunal-vaishnavi Mar 14, 2025
fb18f80
Update MHA op spec again
kunal-vaishnavi Mar 14, 2025
f6aee5f
Update wording in MHA op spec details
kunal-vaishnavi Mar 14, 2025
29396f1
Fix typo in wording
kunal-vaishnavi Mar 14, 2025
1748624
Remove unnecessary commas in op spec
kunal-vaishnavi Mar 14, 2025
edf30d0
Update docs after op spec changes
kunal-vaishnavi Mar 15, 2025
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
26 changes: 17 additions & 9 deletions docs/ContribOperators.md
Original file line number Diff line number Diff line change
Expand Up @@ -1191,17 +1191,17 @@ This version of the operator has been available since version 1 of the 'com.micr
<dd>present state for key with shape (batch_size, num_heads, total_sequence_length, head_size). If past_present_share_buffer is set, its shape is (batch_size, num_heads, max_sequence_length, head_size), while effective_seq_length = (past_sequence_length + kv_sequence_length).</dd>
<dt><tt>present_value</tt> (optional) : T</dt>
<dd>present state for value with shape (batch_size, num_heads, total_sequence_length, head_size). If past_present_share_buffer is set, its shape is (batch_size, num_heads, max_sequence_length, head_size), while effective_seq_length = (past_sequence_length + kv_sequence_length).</dd>
<dt><tt>qk</tt> (optional) : V</dt>
<dt><tt>qk</tt> (optional) : QK</dt>
<dd>normalized Q * K, of shape (batch_size, num_heads, 1, total_sequence_length). </dd>
</dl>

#### Type Constraints

<dl>
<dt><tt>V</tt> : tensor(float)</dt>
<dd>Constrain qk output types to float32 tensors.</dd>
<dt><tt>T</tt> : tensor(float), tensor(float16)</dt>
<dd>Constrain input and output types to float tensors.</dd>
<dt><tt>QK</tt> : tensor(float), tensor(float16)</dt>
<dd>Constrain QK output to float32 or float16 tensors, independent of input type or output type.</dd>
<dt><tt>M</tt> : tensor(int32)</dt>
<dd>Constrain mask index to integer types</dd>
</dl>
Expand Down Expand Up @@ -3203,7 +3203,7 @@ This version of the operator has been available since version 1 of the 'com.micr
<dd>Whether every token can only attend to previous tokens. Default value is 0.</dd>
</dl>

#### Inputs (1 - 8)
#### Inputs (1 - 10)

<dl>
<dt><tt>query</tt> : T</dt>
Expand All @@ -3219,27 +3219,35 @@ This version of the operator has been available since version 1 of the 'com.micr
<dt><tt>attention_bias</tt> (optional) : T</dt>
<dd>bias added to QxK' with shape (batch_size or 1, num_heads or 1, sequence_length, total_sequence_length)</dd>
<dt><tt>past_key</tt> (optional) : T</dt>
<dd>past state for self attention key with shape (batch_size, num_heads, past_sequence_length, head_size)</dd>
<dd>past state for key with shape (batch_size, num_heads, past_sequence_length, head_size) or (batch_size, num_heads, max_sequence_length, head_size) when buffer sharing is used</dd>
<dt><tt>past_value</tt> (optional) : T</dt>
<dd>past state for self attention value with shape (batch_size, num_heads, past_sequence_length, head_size)</dd>
<dd>past state for value with shape (batch_size, num_heads, past_sequence_length, head_size) or (batch_size, num_heads, max_sequence_length, head_size) when buffer sharing is used</dd>
<dt><tt>past_sequence_length</tt> (optional) : M</dt>
<dd>The past_sequence_length buffer sharing is used with</dd>
<dt><tt>cache_indirection</tt> (optional) : M</dt>
<dd>A buffer of shape [batch_size, beam_width, max_sequence_length] where an [i, j, k] entry specifieswhich beam the 'k' th token came from for the 'j' th beam for batch 'i' in the current iteration</dd>
</dl>

#### Outputs (1 - 3)
#### Outputs (1 - 4)

<dl>
<dt><tt>output</tt> : T</dt>
<dd>3D output tensor with shape (batch_size, sequence_length, v_hidden_size)</dd>
<dt><tt>present_key</tt> (optional) : T</dt>
<dd>present state for cross attention key with shape (batch_size, num_heads, kv_sequence_length, head_size)or present state for self attention key with shape (batch_size, num_heads, total_sequence_length, head_size)</dd>
<dd>present state for key with shape (batch_size, num_heads, total_sequence_length, head_size) or (batch_size, num_heads, max_sequence_length, head_size) when buffer sharing is used</dd>
<dt><tt>present_value</tt> (optional) : T</dt>
<dd>present state for cross attention value with shape (batch_size, num_heads, kv_sequence_length, head_size)or present state for self attention value with shape (batch_size, num_heads, total_sequence_length, head_size)</dd>
<dd>present state for value with shape (batch_size, num_heads, total_sequence_length, head_size) or (batch_size, num_heads, max_sequence_length, head_size) when buffer sharing is used</dd>
<dt><tt>qk</tt> (optional) : QK</dt>
<dd>normalized Q * K, of shape (batch_size, num_heads, sequence_length, total_sequence_length). </dd>
</dl>

#### Type Constraints

<dl>
<dt><tt>T</tt> : tensor(float), tensor(float16)</dt>
<dd>Constrain input and output to float tensors.</dd>
<dt><tt>QK</tt> : tensor(float), tensor(float16)</dt>
<dd>Constrain QK output to float32 or float16 tensors, independent of input type or output type.</dd>
<dt><tt>M</tt> : tensor(int32)</dt>
<dd>Constrain mask to integer types</dd>
</dl>
Expand Down
10 changes: 5 additions & 5 deletions docs/OperatorKernels.md
Original file line number Diff line number Diff line change
Expand Up @@ -504,7 +504,7 @@ Do not modify directly.*
|CDist|*in* A:**T**<br> *in* B:**T**<br> *out* C:**T**|1+|**T** = tensor(double), tensor(float)|
|ConvTransposeWithDynamicPads|*in* X:**T**<br> *in* W:**T**<br> *in* Pads:**tensor(int64)**<br> *in* B:**T**<br> *out* Y:**T**|1+|**T** = tensor(float)|
|CropAndResize|*in* X:**T1**<br> *in* rois:**T1**<br> *in* batch_indices:**T2**<br> *in* crop_size:**T2**<br> *out* Y:**T1**|1+|**T1** = tensor(float)<br/> **T2** = tensor(int32)|
|DecoderMaskedMultiHeadAttention|*in* query:**T**<br> *in* key:**T**<br> *in* value:**T**<br> *in* mask_index:**M**<br> *in* attention_bias:**T**<br> *in* past_key:**T**<br> *in* past_value:**T**<br> *in* past_sequence_length:**M**<br> *in* beam_width:**M**<br> *in* cache_indirection:**M**<br> *in* bias:**T**<br> *out* output:**T**<br> *out* present_key:**T**<br> *out* present_value:**T**<br> *out* qk:**V**|1+|**T** = tensor(float)|
|DecoderMaskedMultiHeadAttention|*in* query:**T**<br> *in* key:**T**<br> *in* value:**T**<br> *in* mask_index:**M**<br> *in* attention_bias:**T**<br> *in* past_key:**T**<br> *in* past_value:**T**<br> *in* past_sequence_length:**M**<br> *in* beam_width:**M**<br> *in* cache_indirection:**M**<br> *in* bias:**T**<br> *out* output:**T**<br> *out* present_key:**T**<br> *out* present_value:**T**<br> *out* qk:**QK**|1+|**T** = tensor(float)|
|DequantizeLinear|*in* x:**T1**<br> *in* x_scale:**T2**<br> *in* x_zero_point:**T1**<br> *out* y:**T2**|1+|**T1** = tensor(int16), tensor(int32), tensor(int4), tensor(int8), tensor(uint16), tensor(uint4), tensor(uint8)<br/> **T2** = tensor(float)|
|DynamicQuantizeLSTM|*in* X:**T**<br> *in* W:**T2**<br> *in* R:**T2**<br> *in* B:**T**<br> *in* sequence_lens:**T1**<br> *in* initial_h:**T**<br> *in* initial_c:**T**<br> *in* P:**T**<br> *in* W_scale:**T**<br> *in* W_zero_point:**T2**<br> *in* R_scale:**T**<br> *in* R_zero_point:**T2**<br> *out* Y:**T**<br> *out* Y_h:**T**<br> *out* Y_c:**T**|1+|**T** = tensor(float)<br/> **T1** = tensor(int32)<br/> **T2** = tensor(int8), tensor(uint8)|
|DynamicQuantizeMatMul|*in* A:**T1**<br> *in* B:**T2**<br> *in* b_scale:**T1**<br> *in* b_zero_point:**T2**<br> *in* bias:**T1**<br> *out* Y:**T1**|1+|**T1** = tensor(float)<br/> **T2** = tensor(int8), tensor(uint8)|
Expand All @@ -528,7 +528,7 @@ Do not modify directly.*
|MatMulIntegerToFloat|*in* A:**T1**<br> *in* B:**T2**<br> *in* a_scale:**T3**<br> *in* b_scale:**T3**<br> *in* a_zero_point:**T1**<br> *in* b_zero_point:**T2**<br> *in* bias:**T3**<br> *out* Y:**T3**|1+|**T1** = tensor(int8), tensor(uint8)<br/> **T2** = tensor(int8), tensor(uint8)<br/> **T3** = tensor(float)|
|MatMulNBits|*in* A:**T1**<br> *in* B:**T2**<br> *in* scales:**T1**<br> *in* zero_points:**T3**<br> *in* g_idx:**T4**<br> *in* bias:**T1**<br> *out* Y:**T1**|1+|**T1** = tensor(float), tensor(float16)<br/> **T2** = tensor(uint8)<br/> **T3** = tensor(float), tensor(float16), tensor(uint8)<br/> **T4** = tensor(int32)|
|MaxpoolWithMask|*in* X:**T**<br> *in* M:**tensor(int32)**<br> *out* Y:**T**|1+|**T** = tensor(float)|
|MultiHeadAttention|*in* query:**T**<br> *in* key:**T**<br> *in* value:**T**<br> *in* bias:**T**<br> *in* key_padding_mask:**M**<br> *in* attention_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+|**T** = tensor(float)|
|MultiHeadAttention|*in* query:**T**<br> *in* key:**T**<br> *in* value:**T**<br> *in* bias:**T**<br> *in* key_padding_mask:**M**<br> *in* attention_bias:**T**<br> *in* past_key:**T**<br> *in* past_value:**T**<br> *in* past_sequence_length:**M**<br> *in* cache_indirection:**M**<br> *out* output:**T**<br> *out* present_key:**T**<br> *out* present_value:**T**<br> *out* qk:**QK**|1+|**T** = tensor(float)|
|MurmurHash3|*in* X:**T1**<br> *out* Y:**T2**|1+|**T1** = tensor(double), tensor(float), tensor(int32), tensor(int64), tensor(string), tensor(uint32), tensor(uint64)<br/> **T2** = tensor(int32), tensor(uint32)|
|NGramRepeatBlock|*in* input_ids:**Tid**<br> *in* scores:**T**<br> *out* scores_out:**T**|1+|**T** = tensor(float)<br/> **Tid** = tensor(int64)|
|NhwcMaxPool|*in* x:**T**<br> *out* y:**T**|1+|**T** = tensor(int8), tensor(uint8)|
Expand Down Expand Up @@ -906,7 +906,7 @@ Do not modify directly.*
|ComplexMulConj|*in* A:**T**<br> *in* B:**T**<br> *out* C:**T**|1+|**T** = tensor(float), tensor(float16)|
|ConvTransposeWithDynamicPads|*in* X:**T**<br> *in* W:**T**<br> *in* Pads:**tensor(int64)**<br> *in* B:**T**<br> *out* Y:**T**|1+|**T** = tensor(float)|
|DecoderAttention|*in* query:**T**<br> *in* key:**T**<br> *in* q_weight:**T**<br> *in* kv_weight:**T**<br> *in* bias:**T**<br> *in* key_padding_mask:**B**<br> *in* key_cache:**T**<br> *in* value_cache:**T**<br> *in* static_kv:**B**<br> *in* use_past:**B**<br> *in* has_layer_state:**B**<br> *in* has_key_padding_mask:**B**<br> *out* output:**T**<br> *out* new_key_cache:**T**<br> *out* new_value_cache:**T**|1+|**T** = tensor(float), tensor(float16)|
|DecoderMaskedMultiHeadAttention|*in* query:**T**<br> *in* key:**T**<br> *in* value:**T**<br> *in* mask_index:**M**<br> *in* attention_bias:**T**<br> *in* past_key:**T**<br> *in* past_value:**T**<br> *in* past_sequence_length:**M**<br> *in* beam_width:**M**<br> *in* cache_indirection:**M**<br> *in* bias:**T**<br> *out* output:**T**<br> *out* present_key:**T**<br> *out* present_value:**T**<br> *out* qk:**V**|1+|**T** = tensor(float), tensor(float16)|
|DecoderMaskedMultiHeadAttention|*in* query:**T**<br> *in* key:**T**<br> *in* value:**T**<br> *in* mask_index:**M**<br> *in* attention_bias:**T**<br> *in* past_key:**T**<br> *in* past_value:**T**<br> *in* past_sequence_length:**M**<br> *in* beam_width:**M**<br> *in* cache_indirection:**M**<br> *in* bias:**T**<br> *out* output:**T**<br> *out* present_key:**T**<br> *out* present_value:**T**<br> *out* qk:**QK**|1+|**QK** = tensor(float), tensor(float16)<br/> **T** = tensor(float), tensor(float16)|
|DecoderMaskedSelfAttention|*in* input:**T**<br> *in* weights:**T**<br> *in* bias:**T**<br> *in* mask_index:**M**<br> *in* past:**T**<br> *in* attention_bias:**T**<br> *in* past_sequence_length:**M**<br> *in* beam_width:**M**<br> *in* cache_indirection:**M**<br> *out* output:**T**<br> *out* present:**T**|1+|**T** = tensor(float), tensor(float16)|
|DequantizeLinear|*in* x:**T1**<br> *in* x_scale:**T2**<br> *in* x_zero_point:**T1**<br> *out* y:**T2**|1+|**T1** = tensor(int8), tensor(uint8)<br/> **T2** = tensor(float16)|
|DequantizeWithOrder|*in* input:**Q**<br> *in* scale_input:**S**<br> *out* output:**F**|1+|**F** = tensor(float), tensor(float16)<br/> **Q** = tensor(int8)<br/> **S** = tensor(float)|
Expand All @@ -929,7 +929,7 @@ Do not modify directly.*
|MatMulBnb4|*in* A:**T1**<br> *in* B:**T2**<br> *in* absmax:**T1**<br> *out* Y:**T1**|1+|**T1** = tensor(bfloat16), tensor(float), tensor(float16)<br/> **T2** = tensor(uint8)|
|MatMulNBits|*in* A:**T1**<br> *in* B:**T2**<br> *in* scales:**T1**<br> *in* zero_points:**T3**<br> *in* g_idx:**T4**<br> *in* bias:**T1**<br> *out* Y:**T1**|1+|**T1** = tensor(float), tensor(float16)<br/> **T2** = tensor(uint8)|
|MoE|*in* input:**T**<br> *in* router_probs:**T**<br> *in* fc1_experts_weights:**T**<br> *in* fc1_experts_bias:**T**<br> *in* fc2_experts_weights:**T**<br> *in* fc2_experts_bias:**T**<br> *in* fc3_experts_weights:**T**<br> *in* fc3_experts_bias:**T**<br> *out* output:**T**|1+|**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* attention_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+|**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* attention_bias:**T**<br> *in* past_key:**T**<br> *in* past_value:**T**<br> *in* past_sequence_length:**M**<br> *in* cache_indirection:**M**<br> *out* output:**T**<br> *out* present_key:**T**<br> *out* present_value:**T**<br> *out* qk:**QK**|1+|**QK** = tensor(float), tensor(float16)<br/> **T** = tensor(float), tensor(float16)|
|NGramRepeatBlock|*in* input_ids:**Tid**<br> *in* scores:**T**<br> *out* scores_out:**T**|1+|**T** = tensor(float)<br/> **Tid** = tensor(int64)|
|NhwcConv|*in* X:**T**<br> *in* W:**T**<br> *in* B:**T**<br> *out* Y:**T**|1+|**T** = tensor(float), tensor(float16)|
|PackedAttention|*in* input:**T**<br> *in* weights:**T**<br> *in* bias:**T**<br> *in* token_offset:**M**<br> *in* cumulative_sequence_length:**M**<br> *in* attention_bias:**T**<br> *out* output:**T**|1+|**T** = tensor(float), tensor(float16)|
Expand Down Expand Up @@ -1402,7 +1402,7 @@ Do not modify directly.*
|GroupQueryAttention|*in* query:**T**<br> *in* key:**T**<br> *in* value:**T**<br> *in* past_key:**T**<br> *in* past_value:**T**<br> *in* seqlens_k:**M**<br> *in* total_sequence_length:**M**<br> *in* cos_cache:**T**<br> *in* sin_cache:**T**<br> *in* position_ids:**tensor(int64)**<br> *in* attention_bias:**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)|
|MatMulIntegerToFloat|*in* A:**T1**<br> *in* B:**T2**<br> *in* a_scale:**T3**<br> *in* b_scale:**T3**<br> *in* a_zero_point:**T1**<br> *in* b_zero_point:**T2**<br> *in* bias:**T3**<br> *out* Y:**T3**|1+|**T1** = tensor(int8), tensor(uint8)<br/> **T2** = tensor(int8), tensor(uint8)<br/> **T3** = tensor(float), tensor(float16)|
|MatMulNBits|*in* A:**T1**<br> *in* B:**T2**<br> *in* scales:**T1**<br> *in* zero_points:**T3**<br> *in* g_idx:**T4**<br> *in* bias:**T1**<br> *out* Y:**T1**|1+|**T1** = tensor(float), tensor(float16)<br/> **T2** = tensor(uint8)|
|MultiHeadAttention|*in* query:**T**<br> *in* key:**T**<br> *in* value:**T**<br> *in* bias:**T**<br> *in* key_padding_mask:**M**<br> *in* attention_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)|
|MultiHeadAttention|*in* query:**T**<br> *in* key:**T**<br> *in* value:**T**<br> *in* bias:**T**<br> *in* key_padding_mask:**M**<br> *in* attention_bias:**T**<br> *in* past_key:**T**<br> *in* past_value:**T**<br> *in* past_sequence_length:**M**<br> *in* cache_indirection:**M**<br> *out* output:**T**<br> *out* present_key:**T**<br> *out* present_value:**T**<br> *out* qk:**QK**|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)|
|QAttention|*in* input:**T1**<br> *in* weight:**T2**<br> *in* bias:**T3**<br> *in* input_scale:**T3**<br> *in* weight_scale:**T3**<br> *in* mask_index:**T4**<br> *in* input_zero_point:**T1**<br> *in* weight_zero_point:**T2**<br> *in* past:**T3**<br> *out* output:**T3**<br> *out* present:**T3**|1+|**T1** = tensor(int8), tensor(uint8)<br/> **T2** = tensor(int8), tensor(uint8)<br/> **T3** = tensor(float), tensor(float16)<br/> **T4** = tensor(int32)|
|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)|
Expand Down
2 changes: 1 addition & 1 deletion onnxruntime/contrib_ops/cpu/bert/attention.cc
Original file line number Diff line number Diff line change
Expand Up @@ -335,7 +335,7 @@ Status Attention<T>::Compute(OpKernelContext* context) const {

// Compute the attention score and apply the score to V
return ApplyAttention(Q, K, V, mask_index, past, nullptr /* past_key */, nullptr /* past_value */,
output, nullptr /* present_key */, nullptr /* present_value */,
output, nullptr /* present_key */, nullptr /* present_value */, nullptr /* output_qk */,
batch_size, sequence_length, sequence_length,
parameters.head_size, parameters.v_head_size, parameters.v_hidden_size,
attention_bias, context);
Expand Down
1 change: 1 addition & 0 deletions onnxruntime/contrib_ops/cpu/bert/attention_base.h
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@
#include "core/common/common.h"
#include "core/framework/op_kernel.h"
#include "contrib_ops/cpu/bert/attention_common.h"
#include "contrib_ops/cpu/bert/attention_parameters.h"

namespace onnxruntime {
namespace contrib {
Expand Down
Loading