mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-25 19:48:11 +00:00
Fix bug in GRU when linear_before_reset is true and no bias input is provided (#3797)
* Allocate linear_output_ when linear_before_reset is true and there is no bias input. Add test for this combination.
This commit is contained in:
parent
a24c71af40
commit
b386b41703
3 changed files with 82 additions and 39 deletions
|
|
@ -620,7 +620,11 @@ void UniDirectionalGru<T>::Compute(const gsl::span<const T>& inputs_arg,
|
|||
|
||||
if (linear_before_reset_) {
|
||||
// copy Rbh to linear output
|
||||
gsl::copy(batched_bias_Rh_.subspan(batched_bias_Rh_local - batched_bias_Rh_.begin(), batched_bias_Rh_local_end - batched_bias_Rh_local), linear_output_);
|
||||
if (use_bias_) {
|
||||
gsl::copy(batched_bias_Rh_.subspan(batched_bias_Rh_local - batched_bias_Rh_.begin(),
|
||||
batched_bias_Rh_local_end - batched_bias_Rh_local),
|
||||
linear_output_);
|
||||
}
|
||||
|
||||
// compute Ht-1 * (Rh^T) + Rbh
|
||||
ComputeGemm(batch_size_, hidden_size_, hidden_size_, alpha,
|
||||
|
|
@ -628,7 +632,7 @@ void UniDirectionalGru<T>::Compute(const gsl::span<const T>& inputs_arg,
|
|||
hidden_size_,
|
||||
recurrent_weightsH.cbegin(), recurrent_weightsH.cend(), // Rh^T
|
||||
hidden_size_, beta,
|
||||
linear_output_.begin(), linear_output_.end(), // pre: Rbh, post:output
|
||||
linear_output_.begin(), linear_output_.end(), // pre: Rbh if use_bias_, post:output
|
||||
hidden_size_, ttp_);
|
||||
|
||||
DumpMatrix("Ht-1 * (Rh^T) + Rbh " + seqno_str, linear_output_.data(), batch_size_, hidden_size_);
|
||||
|
|
@ -832,12 +836,17 @@ void UniDirectionalGru<T>::AllocateBuffers() {
|
|||
if (linear_before_reset_) {
|
||||
batched_bias_Wh_ = Allocate(allocator_, batch_size_ * hidden_size_, batched_bias_Wh_ptr_);
|
||||
batched_bias_Rh_ = Allocate(allocator_, batch_size_ * hidden_size_, batched_bias_Rh_ptr_);
|
||||
linear_output_ = Allocate(allocator_, batch_size_ * hidden_size_, linear_output_ptr_);
|
||||
} else {
|
||||
batched_bias_WRh_ = Allocate(allocator_, batch_size_ * hidden_size_, batched_bias_WRh_ptr_);
|
||||
}
|
||||
}
|
||||
|
||||
if (linear_before_reset_) {
|
||||
// if use_bias_ is true we copy bias values to this as the first use. if it's false we don't and need to initialize
|
||||
bool fill = !use_bias_;
|
||||
linear_output_ = Allocate(allocator_, batch_size_ * hidden_size_, linear_output_ptr_, fill);
|
||||
}
|
||||
|
||||
auto batch_times_seq_length = batch_size_ * seq_length_;
|
||||
|
||||
outputZRH_ = Allocate(allocator_, hidden_size_ * 3 * batch_times_seq_length, outputZRH_ptr_, true);
|
||||
|
|
|
|||
|
|
@ -260,9 +260,10 @@ class GRU_ONNXRuntimeUnitTests():
|
|||
print_results(fw_output)
|
||||
|
||||
@staticmethod
|
||||
def BidirectionalDefaultActivationsSimpleWeightsNoBiasTwoRows():
|
||||
def BidirectionalDefaultActivationsSimpleWeightsNoBiasTwoRows(linear_before_reset=0):
|
||||
|
||||
print(GRU_ONNXRuntimeUnitTests.BidirectionalDefaultActivationsSimpleWeightsNoBiasTwoRows.__name__)
|
||||
print(GRU_ONNXRuntimeUnitTests.BidirectionalDefaultActivationsSimpleWeightsNoBiasTwoRows.__name__ +
|
||||
'.linear_before_reset=' + str(linear_before_reset))
|
||||
|
||||
seq_length = 2
|
||||
batch_size = 2
|
||||
|
|
@ -280,7 +281,8 @@ class GRU_ONNXRuntimeUnitTests():
|
|||
gru = GRU_Helper(X=input,
|
||||
W=np.tile(W, (2, 1)).reshape(2, 3 * hidden_size, input_size),
|
||||
R=np.tile(R, (2, 1)).reshape(2, 3 * hidden_size, hidden_size),
|
||||
direction='bidirectional')
|
||||
direction='bidirectional',
|
||||
linear_before_reset=linear_before_reset)
|
||||
|
||||
fw_output = gru.run()
|
||||
print_results(fw_output)
|
||||
|
|
@ -328,7 +330,7 @@ class GRU_ONNXRuntimeUnitTests():
|
|||
GRU_ONNXRuntimeUnitTests.DefaultActivationsSimpleWeightsWithBias(linear_before_reset=1)
|
||||
|
||||
@staticmethod
|
||||
def ReverseDefaultActivationsSimpleWeightsWithBiasBatchParallelLinearBeforeReset(linear_before_reset=0):
|
||||
def ReverseDefaultActivationsSimpleWeightsWithBiasBatchParallelLinearBeforeReset():
|
||||
|
||||
GRU_ONNXRuntimeUnitTests.DefaultActivationsSimpleWeightsWithBias(direction="reverse", linear_before_reset=1)
|
||||
|
||||
|
|
@ -338,7 +340,7 @@ class GRU_ONNXRuntimeUnitTests():
|
|||
GRU_ONNXRuntimeUnitTests.DefaultActivationsSimpleWeightsWithBias(rows=1, linear_before_reset=1)
|
||||
|
||||
@staticmethod
|
||||
def ReverseDefaultActivationsSimpleWeightsWithBiasLinearBeforeReset(linear_before_reset=0):
|
||||
def ReverseDefaultActivationsSimpleWeightsWithBiasLinearBeforeReset():
|
||||
|
||||
GRU_ONNXRuntimeUnitTests.DefaultActivationsSimpleWeightsWithBias(rows=1,
|
||||
direction="reverse",
|
||||
|
|
@ -383,6 +385,7 @@ class GRU_ONNXRuntimeUnitTests():
|
|||
GRU_ONNXRuntimeUnitTests.ForwardDefaultActivationsSimpleWeightsNoBiasTwoRows()
|
||||
GRU_ONNXRuntimeUnitTests.ReverseDefaultActivationsSimpleWeightsNoBiasTwoRows()
|
||||
GRU_ONNXRuntimeUnitTests.BidirectionalDefaultActivationsSimpleWeightsNoBiasTwoRows()
|
||||
GRU_ONNXRuntimeUnitTests.BidirectionalDefaultActivationsSimpleWeightsNoBiasTwoRows(linear_before_reset=1)
|
||||
|
||||
GRU_ONNXRuntimeUnitTests.ForwardDefaultActivationsSimpleWeightsWithBiasBatchParallel()
|
||||
GRU_ONNXRuntimeUnitTests.ForwardDefaultActivationsSimpleWeightsWithBiasBatchParallelLinearBeforeReset()
|
||||
|
|
|
|||
|
|
@ -96,14 +96,15 @@ static void RunGruTest(const std::vector<float>& X_data,
|
|||
} else {
|
||||
test.AddMissingOptionalOutput<float>();
|
||||
}
|
||||
|
||||
|
||||
// TensorRT failed on GRU tests
|
||||
test.Run(OpTester::ExpectResult::kExpectSuccess, "", {kTensorrtExecutionProvider});
|
||||
}
|
||||
|
||||
void DefaultActivationsSimpleWeightsNoBias(std::string direction,
|
||||
const std::vector<float>& Y_data,
|
||||
const std::vector<float>& Y_h_data) {
|
||||
const std::vector<float>& Y_h_data,
|
||||
bool linear_before_reset = false) {
|
||||
int64_t seq_length = 2;
|
||||
int batch_size = 2;
|
||||
int64_t input_size = 1;
|
||||
|
|
@ -127,13 +128,13 @@ void DefaultActivationsSimpleWeightsNoBias(std::string direction,
|
|||
std::vector<float> R_data(num_directions * 3 * hidden_size * hidden_size, 0.1f);
|
||||
|
||||
RunGruTest(X_data, W_data, R_data, Y_data, Y_h_data, input_size, batch_size, hidden_size, seq_length,
|
||||
nullptr, nullptr, nullptr, direction);
|
||||
nullptr, nullptr, nullptr, direction, 9999.0, true, linear_before_reset);
|
||||
|
||||
// if Y_h_data is empty that tests Y_h not being returned. we need to have at least one output or
|
||||
// the node will get removed, so only test with output_sequence == false (no Y as output) if Y_h is not optional
|
||||
if (!Y_h_data.empty())
|
||||
RunGruTest(X_data, W_data, R_data, Y_data, Y_h_data, input_size, batch_size, hidden_size, seq_length,
|
||||
nullptr, nullptr, nullptr, direction, 9999.0, /* output_sequence*/ false);
|
||||
nullptr, nullptr, nullptr, direction, 9999.0, /* output_sequence*/ false, linear_before_reset);
|
||||
}
|
||||
|
||||
TEST(GRUTest, ForwardDefaultActivationsSimpleWeightsNoBiasTwoRows) {
|
||||
|
|
@ -199,6 +200,36 @@ TEST(GRUTest, BidirectionalDefaultActivationsSimpleWeightsNoBiasTwoRows) {
|
|||
DefaultActivationsSimpleWeightsNoBias("bidirectional", Y_data, Y_h_data);
|
||||
}
|
||||
|
||||
TEST(GRUTest, BidirectionalDefaultActivationsSimpleWeightsNoBiasTwoRowsLinearBeforeReset) {
|
||||
std::vector<float> Y_data{
|
||||
// forward output for input sequence 0
|
||||
0.47502081f, 0.450166f, 0.42555748f,
|
||||
0.450166f, 0.40131234f, 0.35434369f,
|
||||
|
||||
// reverse output for input sequence 0 [sequence 1 in reversed input]
|
||||
0.60827853f, 0.50623393f, 0.4426924f,
|
||||
0.5803454f, 0.4527356f, 0.36886264f,
|
||||
|
||||
// forward output for input sequence 1
|
||||
0.60270932f, 0.50830227f, 0.44950222f,
|
||||
0.57543688f, 0.45485455f, 0.37478411f,
|
||||
|
||||
// reverse output for input sequence 1 [sequence 0 in reversed input]
|
||||
0.26894142f, 0.11920292f, 0.04742587f,
|
||||
0.24973989f, 0.09975048f, 0.03557118f};
|
||||
|
||||
std::vector<float> Y_h_data{
|
||||
// we did the forward processing of input[1] last
|
||||
0.60270932f, 0.50830227f, 0.44950222f,
|
||||
0.57543688f, 0.45485455f, 0.37478411f,
|
||||
|
||||
// and the reverse processing of input[0] last as the input order was reversed
|
||||
0.60827853f, 0.50623393f, 0.4426924f,
|
||||
0.5803454f, 0.4527356f, 0.36886264f};
|
||||
|
||||
DefaultActivationsSimpleWeightsNoBias("bidirectional", Y_data, Y_h_data, true);
|
||||
}
|
||||
|
||||
void DefaultActivationsSimpleWeightsWithBias(std::string direction,
|
||||
const std::vector<float>& Y_data,
|
||||
bool linear_before_reset = false,
|
||||
|
|
@ -305,7 +336,7 @@ TEST(GRUTest, ReverseDefaultActivationsSimpleWeightsWithBiasLinearBeforeReset) {
|
|||
}
|
||||
|
||||
/*******************
|
||||
* Tests from ONNXRuntime
|
||||
* Legacy tests from LotusRT
|
||||
*/
|
||||
class DeepCpuGruOpTestContext {
|
||||
public:
|
||||
|
|
@ -464,33 +495,33 @@ void DeepCpuGruOpTestContext::RunTest(const std::vector<float>& X,
|
|||
const std::vector<float>& expected_Y_h,
|
||||
const bool linear_before_reset) {
|
||||
//run with and without output_sequence
|
||||
::onnxruntime::test::RunGruTest(X, gru_input_weights_, gru_recurrent_weights_,
|
||||
expected_Y, expected_Y_h,
|
||||
input_size_, batch_size, hidden_dim_, seq_length,
|
||||
use_bias_ ? &gru_bias_ : nullptr,
|
||||
initial_h,
|
||||
&sequence_lens,
|
||||
direction_,
|
||||
9999999999.f,
|
||||
/*output_sequence*/ true,
|
||||
linear_before_reset,
|
||||
activation_func_names_,
|
||||
alphas_,
|
||||
betas_);
|
||||
RunGruTest(X, gru_input_weights_, gru_recurrent_weights_,
|
||||
expected_Y, expected_Y_h,
|
||||
input_size_, batch_size, hidden_dim_, seq_length,
|
||||
use_bias_ ? &gru_bias_ : nullptr,
|
||||
initial_h,
|
||||
&sequence_lens,
|
||||
direction_,
|
||||
9999999999.f,
|
||||
/*output_sequence*/ true,
|
||||
linear_before_reset,
|
||||
activation_func_names_,
|
||||
alphas_,
|
||||
betas_);
|
||||
|
||||
::onnxruntime::test::RunGruTest(X, gru_input_weights_, gru_recurrent_weights_,
|
||||
expected_Y, expected_Y_h,
|
||||
input_size_, batch_size, hidden_dim_, seq_length,
|
||||
use_bias_ ? &gru_bias_ : nullptr,
|
||||
initial_h,
|
||||
&sequence_lens,
|
||||
direction_,
|
||||
9999999999.f,
|
||||
/*output_sequence*/ false,
|
||||
linear_before_reset,
|
||||
activation_func_names_,
|
||||
alphas_,
|
||||
betas_);
|
||||
RunGruTest(X, gru_input_weights_, gru_recurrent_weights_,
|
||||
expected_Y, expected_Y_h,
|
||||
input_size_, batch_size, hidden_dim_, seq_length,
|
||||
use_bias_ ? &gru_bias_ : nullptr,
|
||||
initial_h,
|
||||
&sequence_lens,
|
||||
direction_,
|
||||
9999999999.f,
|
||||
/*output_sequence*/ false,
|
||||
linear_before_reset,
|
||||
activation_func_names_,
|
||||
alphas_,
|
||||
betas_);
|
||||
}
|
||||
|
||||
TEST(GRUTest, ONNXRuntime_TestGRUOpForwardBasic) {
|
||||
|
|
|
|||
Loading…
Reference in a new issue