mirror of
https://github.com/wassname/keras-js.git
synced 2026-09-09 11:25:25 +08:00
tensor reshaping operation optimizations
This commit is contained in:
@@ -1,8 +1,6 @@
|
||||
import Tensor from '../../Tensor'
|
||||
import Convolution2D from './Convolution2D'
|
||||
import ops from 'ndarray-ops'
|
||||
import unpack from 'ndarray-unpack'
|
||||
import flattenDeep from 'lodash/flattenDeep'
|
||||
|
||||
/**
|
||||
* AtrousConvolution2D layer class
|
||||
@@ -87,18 +85,20 @@ export default class AtrousConvolution2D extends Convolution2D {
|
||||
|
||||
const imColsMat = new Tensor([], [nbPatches, patchLen])
|
||||
|
||||
let patch = new Tensor([], [patchLen])
|
||||
let patch = new Tensor([], [nbRow, nbCol, inputChannels])
|
||||
let patchRaveled = new Tensor([], [patchLen])
|
||||
let n = 0
|
||||
for (let i = 0, limit = inputRows - nbRowDilated; i <= limit; i += this.subsample[0]) {
|
||||
for (let j = 0, limit = inputCols - nbColDilated; j <= limit; j += this.subsample[1]) {
|
||||
const patchData = flattenDeep(unpack(
|
||||
ops.assign(
|
||||
patch.tensor,
|
||||
x.tensor
|
||||
.hi(i + nbRowDilated, j + nbColDilated, inputChannels)
|
||||
.lo(i, j, 0)
|
||||
.step(this.atrousRate[0], this.atrousRate[1], 1)
|
||||
))
|
||||
patch.replaceTensorData(patchData)
|
||||
ops.assign(imColsMat.tensor.pick(n, null), patch.tensor)
|
||||
)
|
||||
patchRaveled.replaceTensorData(patch.tensor.data)
|
||||
ops.assign(imColsMat.tensor.pick(n, null), patchRaveled.tensor)
|
||||
n += 1
|
||||
}
|
||||
}
|
||||
|
||||
@@ -3,8 +3,6 @@ import Tensor from '../../Tensor'
|
||||
import Layer from '../../Layer'
|
||||
import ops from 'ndarray-ops'
|
||||
import gemm from 'ndarray-gemm'
|
||||
import unpack from 'ndarray-unpack'
|
||||
import flattenDeep from 'lodash/flattenDeep'
|
||||
|
||||
/**
|
||||
* Convolution2D layer class
|
||||
@@ -146,15 +144,14 @@ export default class Convolution2D extends Layer {
|
||||
|
||||
const imColsMat = new Tensor([], [nbPatches, patchLen])
|
||||
|
||||
let patch = new Tensor([], [patchLen])
|
||||
let patch = new Tensor([], [nbRow, nbCol, inputChannels])
|
||||
let patchRaveled = new Tensor([], [patchLen])
|
||||
let n = 0
|
||||
for (let i = 0, limit = inputRows - nbRow; i <= limit; i += this.subsample[0]) {
|
||||
for (let j = 0, limit = inputCols - nbCol; j <= limit; j += this.subsample[1]) {
|
||||
const patchData = flattenDeep(unpack(
|
||||
x.tensor.hi(i + nbRow, j + nbCol, inputChannels).lo(i, j, 0)
|
||||
))
|
||||
patch.replaceTensorData(patchData)
|
||||
ops.assign(imColsMat.tensor.pick(n, null), patch.tensor)
|
||||
ops.assign(patch.tensor, x.tensor.hi(i + nbRow, j + nbCol, inputChannels).lo(i, j, 0))
|
||||
patchRaveled.replaceTensorData(patch.tensor.data)
|
||||
ops.assign(imColsMat.tensor.pick(n, null), patchRaveled.tensor)
|
||||
n += 1
|
||||
}
|
||||
}
|
||||
@@ -174,13 +171,12 @@ export default class Convolution2D extends Layer {
|
||||
|
||||
const wRowsMat = new Tensor([], [patchLen, nbFilter])
|
||||
|
||||
let patch = new Tensor([], [patchLen])
|
||||
let patch = new Tensor([], [nbRow, nbCol, inputChannels])
|
||||
let patchRaveled = new Tensor([], [patchLen])
|
||||
for (let n = 0; n < nbFilter; n++) {
|
||||
const patchData = flattenDeep(unpack(
|
||||
this.weights.W.tensor.pick(null, null, null, n)
|
||||
))
|
||||
patch.replaceTensorData(patchData)
|
||||
ops.assign(wRowsMat.tensor.pick(null, n), patch.tensor)
|
||||
ops.assign(patch.tensor, this.weights.W.tensor.pick(null, null, null, n))
|
||||
patchRaveled.replaceTensorData(patch.tensor.data)
|
||||
ops.assign(wRowsMat.tensor.pick(null, n), patchRaveled.tensor)
|
||||
}
|
||||
|
||||
return wRowsMat
|
||||
@@ -197,25 +193,12 @@ export default class Convolution2D extends Layer {
|
||||
x.tensor = x.tensor.transpose(1, 2, 0)
|
||||
}
|
||||
|
||||
let startTime = performance.now()
|
||||
this._calcOutputShape(x)
|
||||
let endTime = performance.now()
|
||||
console.log('_calcOutputShape', endTime - startTime)
|
||||
startTime = performance.now()
|
||||
this._padInput(x)
|
||||
endTime = performance.now()
|
||||
console.log('_padInput', endTime - startTime)
|
||||
|
||||
startTime = performance.now()
|
||||
const imColsMat = this._im2col(x)
|
||||
endTime = performance.now()
|
||||
console.log('imColsMat', endTime - startTime)
|
||||
startTime = performance.now()
|
||||
const wRowsMat = this._w2row(x)
|
||||
endTime = performance.now()
|
||||
console.log('wRowsMat', endTime - startTime)
|
||||
|
||||
startTime = performance.now()
|
||||
const nbFilter = this.kernelShape[0]
|
||||
const outputRows = this.outputShape[0]
|
||||
const outputCols = this.outputShape[1]
|
||||
@@ -226,10 +209,7 @@ export default class Convolution2D extends Layer {
|
||||
ops.assigns(matMul.tensor.pick(null, n), this.weights.b.tensor.get(n))
|
||||
}
|
||||
}
|
||||
endTime = performance.now()
|
||||
console.log('createMatMul', endTime - startTime)
|
||||
|
||||
startTime = performance.now()
|
||||
if (x._useWeblas) {
|
||||
const bias = this.bias
|
||||
? this.weights.b.tensor.data
|
||||
@@ -242,21 +222,15 @@ export default class Convolution2D extends Layer {
|
||||
} else {
|
||||
gemm(matMul.tensor, imColsMat.tensor, wRowsMat.tensor, 1, 1)
|
||||
}
|
||||
endTime = performance.now()
|
||||
console.log('gemm', endTime - startTime)
|
||||
|
||||
startTime = performance.now()
|
||||
let output = new Tensor([], this.outputShape)
|
||||
let outputChannelRaveled = new Tensor([], [outputRows * outputCols])
|
||||
let outputChannel = new Tensor([], [outputRows, outputCols])
|
||||
for (let n = 0; n < nbFilter; n++) {
|
||||
const outputChannelData = flattenDeep(unpack(
|
||||
matMul.tensor.pick(null, n)
|
||||
))
|
||||
outputChannel.replaceTensorData(outputChannelData)
|
||||
ops.assign(outputChannelRaveled.tensor, matMul.tensor.pick(null, n))
|
||||
outputChannel.replaceTensorData(outputChannelRaveled.tensor.data)
|
||||
ops.assign(output.tensor.pick(null, null, n), outputChannel.tensor)
|
||||
}
|
||||
endTime = performance.now()
|
||||
console.log('createOutput', endTime - startTime)
|
||||
x.tensor = output.tensor
|
||||
|
||||
this.activation(x)
|
||||
|
||||
@@ -3,8 +3,6 @@ import Tensor from '../../Tensor'
|
||||
import Layer from '../../Layer'
|
||||
import ops from 'ndarray-ops'
|
||||
import gemm from 'ndarray-gemm'
|
||||
import unpack from 'ndarray-unpack'
|
||||
import flattenDeep from 'lodash/flattenDeep'
|
||||
|
||||
/**
|
||||
* Convolution3D layer class
|
||||
@@ -163,16 +161,15 @@ export default class Convolution3D extends Layer {
|
||||
|
||||
const volColsMat = new Tensor([], [nbPatches, patchLen])
|
||||
|
||||
let patch = new Tensor([], [patchLen])
|
||||
let patch = new Tensor([], [kernelDim1, kernelDim2, kernelDim3, inputChannels])
|
||||
let patchRaveled = new Tensor([], [patchLen])
|
||||
let n = 0
|
||||
for (let i = 0, limit = inputDim1 - kernelDim1; i <= limit; i += this.subsample[0]) {
|
||||
for (let j = 0, limit = inputDim2 - kernelDim2; j <= limit; j += this.subsample[1]) {
|
||||
for (let k = 0, limit = inputDim3 - kernelDim3; k <= limit; k += this.subsample[2]) {
|
||||
const patchData = flattenDeep(unpack(
|
||||
x.tensor.hi(i + kernelDim1, j + kernelDim2, k + kernelDim3, inputChannels).lo(i, j, k, 0)
|
||||
))
|
||||
patch.replaceTensorData(patchData)
|
||||
ops.assign(volColsMat.tensor.pick(n, null), patch.tensor)
|
||||
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)
|
||||
n += 1
|
||||
}
|
||||
}
|
||||
@@ -193,13 +190,12 @@ export default class Convolution3D extends Layer {
|
||||
|
||||
const wRowsMat = new Tensor([], [patchLen, nbFilter])
|
||||
|
||||
let patch = new Tensor([], [patchLen])
|
||||
let patch = new Tensor([], [kernelDim1, kernelDim2, kernelDim3, inputChannels])
|
||||
let patchRaveled = new Tensor([], [patchLen])
|
||||
for (let n = 0; n < nbFilter; n++) {
|
||||
const patchData = flattenDeep(unpack(
|
||||
this.weights.W.tensor.pick(null, null, null, null, n)
|
||||
))
|
||||
patch.replaceTensorData(patchData)
|
||||
ops.assign(wRowsMat.tensor.pick(null, n), patch.tensor)
|
||||
ops.assign(patch.tensor, this.weights.W.tensor.pick(null, null, null, null, n))
|
||||
patchRaveled.replaceTensorData(patch.tensor.data)
|
||||
ops.assign(wRowsMat.tensor.pick(null, n), patchRaveled.tensor)
|
||||
}
|
||||
|
||||
return wRowsMat
|
||||
@@ -248,12 +244,11 @@ export default class Convolution3D extends Layer {
|
||||
}
|
||||
|
||||
let output = new Tensor([], this.outputShape)
|
||||
let outputChannelRaveled = new Tensor([], [outputDim1 * outputDim2 * outputDim3])
|
||||
let outputChannel = new Tensor([], [outputDim1, outputDim2, outputDim3])
|
||||
for (let n = 0; n < nbFilter; n++) {
|
||||
const outputChannelData = flattenDeep(unpack(
|
||||
matMul.tensor.pick(null, n)
|
||||
))
|
||||
outputChannel.replaceTensorData(outputChannelData)
|
||||
ops.assign(outputChannelRaveled.tensor, matMul.tensor.pick(null, n))
|
||||
outputChannel.replaceTensorData(outputChannelRaveled.tensor.data)
|
||||
ops.assign(output.tensor.pick(null, null, null, n), outputChannel.tensor)
|
||||
}
|
||||
x.tensor = output.tensor
|
||||
|
||||
@@ -3,8 +3,6 @@ import Tensor from '../../Tensor'
|
||||
import Layer from '../../Layer'
|
||||
import ops from 'ndarray-ops'
|
||||
import gemm from 'ndarray-gemm'
|
||||
import unpack from 'ndarray-unpack'
|
||||
import flattenDeep from 'lodash/flattenDeep'
|
||||
|
||||
/**
|
||||
* Deconvolution2D layer class
|
||||
@@ -129,13 +127,12 @@ export default class Deconvolution2D extends Layer {
|
||||
const [inputRows, inputCols, inputChannels] = x.tensor.shape
|
||||
|
||||
const imColsMat = new Tensor([], [inputRows * inputCols, inputChannels])
|
||||
let channel = new Tensor([], [inputRows * inputCols])
|
||||
let channelRaveled = new Tensor([], [inputRows * inputCols])
|
||||
let channel = new Tensor([], [inputRows, inputCols])
|
||||
for (let c = 0; c < inputChannels; c++) {
|
||||
const channelData = flattenDeep(unpack(
|
||||
x.tensor.pick(null, null, c)
|
||||
))
|
||||
channel.replaceTensorData(channelData)
|
||||
ops.assign(imColsMat.tensor.pick(null, c), channel.tensor)
|
||||
ops.assign(channel.tensor, x.tensor.pick(null, null, c))
|
||||
channelRaveled.replaceTensorData(channel.tensor.data)
|
||||
ops.assign(imColsMat.tensor.pick(null, c), channelRaveled.tensor)
|
||||
}
|
||||
return imColsMat
|
||||
}
|
||||
@@ -150,13 +147,12 @@ export default class Deconvolution2D extends Layer {
|
||||
const [nbRow, nbCol, inputChannels, nbFilter] = this.weights.W.tensor.shape
|
||||
|
||||
const wRowsMat = new Tensor([], [inputChannels, nbRow * nbCol * nbFilter])
|
||||
let channel = new Tensor([], [nbRow * nbCol * nbFilter])
|
||||
let channelRaveled = new Tensor([], [nbRow * nbCol * nbFilter])
|
||||
let channel = new Tensor([], [nbRow, nbCol, nbFilter])
|
||||
for (let c = 0; c < inputChannels; c++) {
|
||||
const channelData = flattenDeep(unpack(
|
||||
this.weights.W.tensor.pick(null, null, c, null)
|
||||
))
|
||||
channel.replaceTensorData(channelData)
|
||||
ops.assign(wRowsMat.tensor.pick(c, null), channel.tensor)
|
||||
ops.assign(channel.tensor, this.weights.W.tensor.pick(null, null, c, null))
|
||||
channelRaveled.replaceTensorData(channel.tensor.data)
|
||||
ops.assign(wRowsMat.tensor.pick(c, null), channelRaveled.tensor)
|
||||
}
|
||||
return wRowsMat
|
||||
}
|
||||
@@ -211,11 +207,12 @@ export default class Deconvolution2D extends Layer {
|
||||
|
||||
const patchShape = [nbRow, nbCol, nbFilter]
|
||||
let patch = new Tensor([], patchShape)
|
||||
let patchRaveled = new Tensor([], [nbRow * nbCol * nbFilter])
|
||||
let index = 0
|
||||
for (let i = 0; i < inputRows; i++) {
|
||||
for (let j = 0; j < inputCols; j++) {
|
||||
const patchData = unpack(matMul.tensor.pick(index, null))
|
||||
patch.replaceTensorData(patchData)
|
||||
ops.assign(patchRaveled.tensor, matMul.tensor.pick(index, null))
|
||||
patch.replaceTensorData(patchRaveled.tensor.data)
|
||||
const iOutPos = i * this.subsample[0]
|
||||
const jOutPos = j * this.subsample[1]
|
||||
ops.addeq(
|
||||
|
||||
@@ -1,12 +1,9 @@
|
||||
import Tensor from '../../Tensor'
|
||||
import Layer from '../../Layer'
|
||||
import ndarray from 'ndarray'
|
||||
import unpack from 'ndarray-unpack'
|
||||
import flattenDeep from 'lodash/flattenDeep'
|
||||
|
||||
/**
|
||||
* Flatten layer class
|
||||
* Turns tensor into 1-d. Note there is no concept of batch size in these layers (single-batch).
|
||||
* We use ndarray-unpack first, as ndarray striding/offsets precludes us from simply using x.tensor.data
|
||||
*/
|
||||
export default class Flatten extends Layer {
|
||||
/**
|
||||
@@ -24,8 +21,9 @@ export default class Flatten extends Layer {
|
||||
*/
|
||||
call (x) {
|
||||
if (x.tensor.shape.length > 1) {
|
||||
const shape = [x.tensor.shape.reduce((a, b) => a * b, 1)]
|
||||
x.tensor = ndarray(new x._type(flattenDeep(unpack(x.tensor))), shape)
|
||||
let raveled = new Tensor([], [x.tensor.shape.reduce((a, b) => a * b, 1)])
|
||||
raveled.replaceTensorData(x.tensor.data)
|
||||
x.tensor = raveled.tensor
|
||||
}
|
||||
return x
|
||||
}
|
||||
|
||||
@@ -1,12 +1,9 @@
|
||||
import Tensor from '../../Tensor'
|
||||
import Layer from '../../Layer'
|
||||
import ndarray from 'ndarray'
|
||||
import unpack from 'ndarray-unpack'
|
||||
import flattenDeep from 'lodash/flattenDeep'
|
||||
|
||||
/**
|
||||
* Reshape layer class
|
||||
* Note there is no concept of batch size in these layers (single-batch).
|
||||
* We use ndarray-unpack first, as ndarray striding/offsets precludes us from simply using x.tensor.data
|
||||
*/
|
||||
export default class Reshape extends Layer {
|
||||
/**
|
||||
@@ -32,7 +29,9 @@ export default class Reshape extends Layer {
|
||||
if (this.targetShape.reduce((a, b) => a * b, 1) !== x.tensor.size) {
|
||||
throw new Error(`${this.name} [Reshape layer] The total size of new array must be unchanged in reshape layer.`)
|
||||
}
|
||||
x.tensor = ndarray(new x._type(flattenDeep(unpack(x.tensor))), this.targetShape)
|
||||
let reshaped = new Tensor([], this.targetShape)
|
||||
reshaped.replaceTensorData(x.tensor.data)
|
||||
x.tensor = reshaped.tensor
|
||||
return x
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user