Welcome to mirror list, hosted at ThFree Co, Russian Federation.

github.com/torch/nn.git - Unnamed repository; edit this file 'description' to name the repository.
summaryrefslogtreecommitdiff
path: root/test
diff options
context:
space:
mode:
authornicholas-leonard <nick@nikopia.org>2014-05-13 08:13:12 +0400
committernicholas-leonard <nick@nikopia.org>2014-05-13 08:13:12 +0400
commiteaa42b957a79e58071762ac44d48189ff7ed013b (patch)
tree4da72bdd41c6037164272df7caf048194a65142b /test
parent0042b6138cff3e0f6092c4e6c72517b8927a3d7b (diff)
modified TemporalMaxPooling unit test to capture corner case
Diffstat (limited to 'test')
-rw-r--r--test/test.lua8
1 files changed, 4 insertions, 4 deletions
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()