mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-19 19:00:47 +00:00
Fix mask in EmbedLayerNormalization (#2300)
This commit is contained in:
parent
6e65dcf588
commit
a6b2c9fc09
2 changed files with 82 additions and 1 deletions
|
|
@ -182,7 +182,7 @@ bool LaunchEmbedLayerNormKernel(
|
|||
const size_t element_size) {
|
||||
const cudaStream_t stream = nullptr; // default stream
|
||||
|
||||
if (!ComputeMaskIndex(stream, sequence_length, hidden_size, input_mask, static_cast<int*>(mask_index))) {
|
||||
if (!ComputeMaskIndex(stream, sequence_length, batch_size, input_mask, static_cast<int*>(mask_index))) {
|
||||
return false;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -269,5 +269,86 @@ TEST(EmbedLayerNormTest, EmbedLayerNormBatch2) {
|
|||
sequence_length,
|
||||
hidden_size);
|
||||
}
|
||||
|
||||
// BatchSize > HiddenSize to reproduce mask processing bug
|
||||
TEST(EmbedLayerNormTest, EmbedLayerNormLargeBatchSmallHiddenSize) {
|
||||
int batch_size = 5;
|
||||
int sequence_length = 2;
|
||||
int hidden_size = 4;
|
||||
|
||||
std::vector<int32_t> input_ids_data = {
|
||||
1, 3,
|
||||
1, 3,
|
||||
2, 0,
|
||||
1, 3,
|
||||
2, 0};
|
||||
|
||||
std::vector<int32_t> segment_ids_data = {
|
||||
0, 1,
|
||||
0, 1,
|
||||
0, 0,
|
||||
0, 1,
|
||||
0, 0};
|
||||
|
||||
std::vector<int32_t> mask_data = {
|
||||
1, 1,
|
||||
1, 1,
|
||||
1, 0,
|
||||
1, 1,
|
||||
1, 0};
|
||||
|
||||
std::vector<float> word_embedding_data = {
|
||||
0.2f, 0.1f, 0.4f, -0.6f,
|
||||
0.3f, 0.2f, 0.5f, 0.6f,
|
||||
0.6f, 0.7f, 0.0f, -0.1f,
|
||||
0.8f, 0.6f, 0.9f, 1.2f,
|
||||
0.1f, 0.3f, 0.5f, 0.9f,
|
||||
1.0f, -2.0f, 1.1f, 0.8f};
|
||||
|
||||
std::vector<float> position_embedding_data = {
|
||||
0.1f, 0.1f, 0.4f, 0.6f,
|
||||
0.6f, 0.0f, 0.8f, 0.6f,
|
||||
0.3f, 0.9f, -2.0f, 0.8f};
|
||||
|
||||
std::vector<float> segment_embedding_data = {
|
||||
0.3f, 0.4f, 0.9f, 0.1f,
|
||||
0.7f, 0.3f, 0.5f, 0.2f};
|
||||
|
||||
std::vector<float> gamma_data = {
|
||||
0.25f, 0.15f, 0.45f, -0.66f};
|
||||
|
||||
std::vector<float> beta_data = {
|
||||
0.6f, 0.2f, 0.5f, -0.6f};
|
||||
|
||||
std::vector<float> output_data = {
|
||||
0.36917170882225037, 0.061503000557422638, 1.1598974466323853, -0.85092413425445557,
|
||||
0.74301940202713013, -0.057434864342212677, 0.84324657917022705, -0.85171419382095337,
|
||||
0.36917170882225037, 0.061503000557422638, 1.1598974466323853, -0.85092413425445557,
|
||||
0.74301940202713013, -0.057434864342212677, 0.84324657917022705, -0.85171419382095337,
|
||||
0.57668739557266235, 0.2979130744934082, 0.96158987283706665, 0.44627034664154053,
|
||||
0.64977931976318359, 0.11039737612009048, 1.1869535446166992, 0.14469735324382782,
|
||||
0.36917170882225037, 0.061503000557422638, 1.1598974466323853, -0.85092413425445557,
|
||||
0.74301940202713013, -0.057434864342212677, 0.84324657917022705, -0.85171419382095337,
|
||||
0.57668739557266235, 0.2979130744934082, 0.96158987283706665, 0.44627034664154053,
|
||||
0.64977931976318359, 0.11039737612009048, 1.1869535446166992, 0.14469735324382782
|
||||
};
|
||||
|
||||
std::vector<int32_t> mask_index_data = {
|
||||
2, 2, 1, 2, 1};
|
||||
|
||||
RunTest(input_ids_data,
|
||||
segment_ids_data,
|
||||
mask_data,
|
||||
word_embedding_data,
|
||||
position_embedding_data,
|
||||
segment_embedding_data,
|
||||
gamma_data,
|
||||
beta_data,
|
||||
output_data,
|
||||
mask_index_data,
|
||||
batch_size,
|
||||
sequence_length,
|
||||
hidden_size);
|
||||
}
|
||||
} // namespace test
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
Loading…
Reference in a new issue