mirror of
https://github.com/wassname/keras-js.git
synced 2026-09-11 12:20:53 +08:00
update Convolution3D and AtrousConvolution2D layers
This commit is contained in:
@@ -83,7 +83,9 @@ export default class AtrousConvolution2D extends Convolution2D {
|
||||
const nbRowDilated = nbRow + (nbRow - 1) * (this.atrousRate[0] - 1)
|
||||
const nbColDilated = nbCol + (nbCol - 1) * (this.atrousRate[1] - 1)
|
||||
|
||||
const imColsMat = new Tensor([], [nbPatches, patchLen])
|
||||
if (!this._imColsMat) {
|
||||
this._imColsMat = new Tensor([], [nbPatches, patchLen])
|
||||
}
|
||||
|
||||
let patch = new Tensor([], [nbRow, nbCol, inputChannels])
|
||||
let patchRaveled = new Tensor([], [patchLen])
|
||||
@@ -98,11 +100,13 @@ export default class AtrousConvolution2D extends Convolution2D {
|
||||
.step(this.atrousRate[0], this.atrousRate[1], 1)
|
||||
)
|
||||
patchRaveled.replaceTensorData(patch.tensor.data)
|
||||
ops.assign(imColsMat.tensor.pick(n, null), patchRaveled.tensor)
|
||||
ops.assign(this._imColsMat.tensor.pick(n, null), patchRaveled.tensor)
|
||||
n += 1
|
||||
}
|
||||
}
|
||||
|
||||
return imColsMat
|
||||
if (this._useWeblas) {
|
||||
this._imColsMat.createWeblasTensor()
|
||||
}
|
||||
return this._imColsMat
|
||||
}
|
||||
}
|
||||
|
||||
@@ -171,7 +171,20 @@ export default class Convolution3D extends Layer {
|
||||
const nbPatches = outputDim1 * outputDim2 * outputDim3
|
||||
const patchLen = kernelDim1 * kernelDim2 * kernelDim3 * inputChannels
|
||||
|
||||
const volColsMat = new Tensor([], [nbPatches, patchLen])
|
||||
if (!this._volColsMat) {
|
||||
this._volColsMat = new Tensor([], [nbPatches, patchLen])
|
||||
}
|
||||
|
||||
if (
|
||||
kernelDim1 === 1 && kernelDim2 === 1 && kernelDim3 === 1 &&
|
||||
this.subsample[0] === 1 && this.subsample[1] === 1 && this.subsample[2] === 1
|
||||
) {
|
||||
this._volColsMat.replaceTensorData(x.tensor.data)
|
||||
if (this._useWeblas) {
|
||||
this._volColsMat.createWeblasTensor()
|
||||
}
|
||||
return this._volColsMat
|
||||
}
|
||||
|
||||
let patch = new Tensor([], [kernelDim1, kernelDim2, kernelDim3, inputChannels])
|
||||
let patchRaveled = new Tensor([], [patchLen])
|
||||
@@ -181,13 +194,15 @@ export default class Convolution3D extends Layer {
|
||||
for (let k = 0, limit = inputDim3 - kernelDim3; k <= limit; k += this.subsample[2]) {
|
||||
ops.assign(patch.tensor, x.tensor.hi(i + kernelDim1, j + kernelDim2, k + kernelDim3, inputChannels).lo(i, j, k, 0))
|
||||
patchRaveled.replaceTensorData(patch.tensor.data)
|
||||
ops.assign(volColsMat.tensor.pick(n, null), patchRaveled.tensor)
|
||||
ops.assign(this._volColsMat.tensor.pick(n, null), patchRaveled.tensor)
|
||||
n += 1
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return volColsMat
|
||||
if (this._useWeblas) {
|
||||
this._volColsMat.createWeblasTensor()
|
||||
}
|
||||
return this._volColsMat
|
||||
}
|
||||
|
||||
/**
|
||||
@@ -226,10 +241,7 @@ export default class Convolution3D extends Layer {
|
||||
this._calcOutputShape(x)
|
||||
this._padInput(x)
|
||||
|
||||
const volColsMat = this._vol2col(x)
|
||||
if (this._useWeblas) {
|
||||
volColsMat.createWeblasTensor()
|
||||
}
|
||||
this._vol2col(x)
|
||||
|
||||
const nbFilter = this.kernelShape[0]
|
||||
const outputDim1 = this.outputShape[0]
|
||||
@@ -241,18 +253,16 @@ export default class Convolution3D extends Layer {
|
||||
if (this._useWeblas) {
|
||||
const bias = this.bias ? this.weights.b.weblasTensor : this._zerosVec.weblasTensor
|
||||
matMul.tensor.data = weblas.pipeline.sgemm(
|
||||
1, volColsMat.weblasTensor, this._wRowsMat.weblasTensor,
|
||||
1, this._volColsMat.weblasTensor, this._wRowsMat.weblasTensor,
|
||||
1, bias
|
||||
).transfer()
|
||||
volColsMat.weblasTensor.delete()
|
||||
delete volColsMat.weblasTensor
|
||||
} else {
|
||||
if (this.bias) {
|
||||
for (let n = 0; n < nbFilter; n++) {
|
||||
ops.assigns(matMul.tensor.pick(null, n), this.weights.b.tensor.get(n))
|
||||
}
|
||||
}
|
||||
gemm(matMul.tensor, volColsMat.tensor, this._wRowsMat.tensor, 1, 1)
|
||||
gemm(matMul.tensor, this._volColsMat.tensor, this._wRowsMat.tensor, 1, 1)
|
||||
}
|
||||
|
||||
let output = new Tensor([], this.outputShape)
|
||||
|
||||
Reference in New Issue
Block a user