From eaa42b957a79e58071762ac44d48189ff7ed013b Mon Sep 17 00:00:00 2001 From: nicholas-leonard Date: Tue, 13 May 2014 00:13:12 -0400 Subject: modified TemporalMaxPooling unit test to capture corner case --- test/test.lua | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) (limited to 'test') diff --git a/test/test.lua b/test/test.lua index ff4e649..9b2a2c9 100644 --- a/test/test.lua +++ b/test/test.lua @@ -1468,13 +1468,13 @@ function nntest.TemporalMaxPooling() local outputGrad = torch.randn(output:size()) local inputGrad = module:backward(input, outputGrad):clone() - local input1D = input:select(1, 1) + local input1D = input:select(1, 2) local output1D = module:forward(input1D) - local outputGrad1D = outputGrad:select(1, 1) + local outputGrad1D = outputGrad:select(1, 2) local inputGrad1D = module:backward(input1D, outputGrad1D) - mytester:assertTensorEq(output:select(1,1), output1D, 0.000001, 'error on 2D vs 1D forward)') - mytester:assertTensorEq(inputGrad:select(1,1), inputGrad1D, 0.000001, 'error on 2D vs 1D backward)') + mytester:assertTensorEq(output:select(1,2), output1D, 0.000001, 'error on 2D vs 1D forward)') + mytester:assertTensorEq(inputGrad:select(1,2), inputGrad1D, 0.000001, 'error on 2D vs 1D backward)') end function nntest.VolumetricConvolution() -- cgit v1.2.3