support th dim ordering in Convolution2D

This commit is contained in:
Leon Chen
2016-08-26 18:56:59 -04:00
parent 3ec98ef701
commit 263242e6e4
+27 -3
View File
@@ -39,10 +39,10 @@ export default class Convolution2D extends Layer {
this.subsample = subsample
if (dimOrdering !== 'tf') {
throw new Error(`${this.name} [Convolution2D layer] Only tf dim ordering supported currently.`)
} else {
if (dimOrdering === 'tf' || dimOrdering === 'th') {
this.dimOrdering = dimOrdering
} else {
throw new Error(`${this.name} [Convolution2D layer] Only tf and th dim ordering are allowed.`)
}
this.bias = bias
@@ -51,6 +51,25 @@ export default class Convolution2D extends Layer {
this.params = this.bias ? ['W', 'b'] : ['W']
}
/**
* Method for setting layer weights. Extends `super` method.
* W weight tensor is converted to `tf` mode if in `th` mode.
* In `tf` mode, W weight tensor has shape [nbRow, nbCol, inputChannels, nbFilter]
* In `th` mode, W weight tensor has shape [nbFilter, inputChannels, nbRow, nbCol]
* @param {Tensor[]} weightsArr - array of weights which are instances of Tensor
*/
setWeights = weightsArr => {
if (this.dimOrdering === 'th') {
const weightsArrTheano = weightsArr.map(w => {
w.tensor = w.tensor.transpose(3, 2, 0, 1)
return w
})
super.setWeights(weightsArrTheano)
} else {
super.setWeights(weightsArr)
}
}
/**
* Method for computing output dimensions and padding, based on input
* dimensions, kernel size, and padding mode.
@@ -172,6 +191,11 @@ export default class Convolution2D extends Layer {
* @returns {Tensor} x
*/
call = x => {
// convert to tf ordering
if (this.dimOrdering === 'th') {
x.tensor = x.tensor.transpose(2, 0, 1)
}
this._calcOutputShape(x)
this._padInput(x)