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-10 23:44:16 +0400
committernicholas-leonard <nick@nikopia.org>2014-05-10 23:44:16 +0400
commit15f3d6b12cf27a14deeab92c4bcb7bc2fb416149 (patch)
tree3ee9dcfc93bb27317e587e1b1ea9c195b84580f5 /test
parent5703059443c6f5a5bfdf5c6ab035d2a377a821e5 (diff)
TemporalConvolution unit test 1D vs 2D
Diffstat (limited to 'test')
-rw-r--r--test/test.lua13
1 files changed, 13 insertions, 0 deletions
diff --git a/test/test.lua b/test/test.lua
index f565bd3..10867fe 100644
--- a/test/test.lua
+++ b/test/test.lua
@@ -1350,6 +1350,19 @@ function nntest.TemporalConvolution()
local ferr, berr = jac.testIO(module, input)
mytester:asserteq(0, ferr, torch.typename(module) .. ' - i/o forward err ')
mytester:asserteq(0, berr, torch.typename(module) .. ' - i/o backward err ')
+
+ -- 2D matches 1D
+ local output = module:forward(input)
+ local outputGrad = torch.randn(output:size())
+ local inputGrad = module:backward(input, outputGrad)
+
+ local input1D = input:select(1, 1)
+ local output1D = module:forward(input1D)
+ local outputGrad1D = outputGrad:select(1, 1)
+ 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)')
end
function nntest.TemporalSubSampling()