diff --git a/demos/src/mnist-cnn.css b/demos/src/mnist-cnn.css
index 5a13811..8be428d 100644
--- a/demos/src/mnist-cnn.css
+++ b/demos/src/mnist-cnn.css
@@ -91,3 +91,43 @@
.demo.mnist-cnn .output .output-class.predicted .output-label {
color: #1BBC9B;
}
+
+.demo.mnist-cnn .layer-result-container {
+ position: relative;
+}
+
+.demo.mnist-cnn .layer-result-container .bg-line {
+ position: absolute;
+ z-index: 0;
+ top: 0;
+ left: 50%;
+ background: white;
+ width: 10px;
+ height: 100%;
+}
+
+.demo.mnist-cnn .layer-result {
+ position: relative;
+ z-index: 1;
+ border-radius: 10px;
+ background: white;
+ margin: 30px 20px;
+ padding: 20px;
+}
+
+.demo.mnist-cnn .layer-result-heading {
+ font-size: 1rem;
+ text-transform: uppercase;
+ margin-bottom: 10px;
+}
+
+.demo.mnist-cnn .layer-result-canvas-container {
+ display: inline-flex;
+ flex-wrap: wrap;
+ background: white;
+}
+
+.demo.mnist-cnn .layer-result-canvas-container canvas {
+ border: 1px solid lightgray;
+ margin: 1px;
+}
diff --git a/demos/src/mnist-cnn.js b/demos/src/mnist-cnn.js
index 1eba47f..8e48a75 100644
--- a/demos/src/mnist-cnn.js
+++ b/demos/src/mnist-cnn.js
@@ -63,8 +63,27 @@ export const MnistCnn = Vue.extend({
-
-
+
+
+
+
{{ layerResult.name }}
+
+
+
+
+
`,
@@ -76,6 +95,7 @@ export const MnistCnn = Vue.extend({
input: new Float32Array(784),
output: new Float32Array(10),
outputClasses: range(10),
+ layerResultImages: [],
drawing: false,
strokes: []
}
@@ -104,6 +124,7 @@ export const MnistCnn = Vue.extend({
methods: {
clear: function (e) {
+ this.clearIntermediateResults()
const ctx = document.getElementById('input-canvas').getContext('2d')
ctx.clearRect(0, 0, ctx.canvas.width, ctx.canvas.height)
const ctxCenterCrop = document.getElementById('input-canvas-centercrop').getContext('2d')
@@ -187,7 +208,56 @@ export const MnistCnn = Vue.extend({
}
this.output = this.model.predict({ input: this.input }).output
- }, 200, { leading: true, trailing: true })
+ this.getIntermediateResults()
+ }, 200, { leading: true, trailing: true }),
+ getIntermediateResults: function () {
+ const layersToShow = ['convolution2d_1', 'convolution2d_2', 'maxpooling2d_1']
+ let results = []
+ layersToShow.forEach(name => {
+ const layer = this.model.modelLayersMap.get(name)
+ let images = []
+ if (layer.result.tensor.shape.length === 3) {
+ images = utils.unroll3Dtensor(layer.result.tensor)
+ } else if (layer.result.tensor.shape.length === 2) {
+ images = [utils.image2Dtensor(layer.result.tensor)]
+ }
+ results.push({
+ name,
+ images
+ })
+ })
+ this.layerResultImages = results
+ setTimeout(() => {
+ this.showIntermediateResults()
+ }, 0)
+ },
+
+ showIntermediateResults: function () {
+ this.layerResultImages.forEach((result, layerNum) => {
+ result.images.forEach((image, imageNum) => {
+ let ctx = document.getElementById(`intermediate-result-${layerNum}-${imageNum}`).getContext('2d')
+ ctx.putImageData(image, 0, 0)
+ let ctxScaled = document.getElementById(`intermediate-result-${layerNum}-${imageNum}-scaled`).getContext('2d')
+ ctxScaled.save()
+ ctxScaled.scale(3, 3)
+ ctxScaled.clearRect(0, 0, ctxScaled.canvas.width, ctxScaled.canvas.height)
+ ctxScaled.drawImage(document.getElementById(`intermediate-result-${layerNum}-${imageNum}`), 0, 0)
+ ctxScaled.restore()
+ })
+ })
+ },
+
+ clearIntermediateResults: function () {
+ this.layerResultImages.forEach((result, layerNum) => {
+ result.images.forEach((image, imageNum) => {
+ let ctxScaled = document.getElementById(`intermediate-result-${layerNum}-${imageNum}-scaled`).getContext('2d')
+ ctxScaled.save()
+ ctxScaled.scale(3, 3)
+ ctxScaled.clearRect(0, 0, ctxScaled.canvas.width, ctxScaled.canvas.height)
+ ctxScaled.restore()
+ })
+ })
+ }
}
})
diff --git a/demos/src/utils/index.js b/demos/src/utils/index.js
index 063cb83..7804213 100644
--- a/demos/src/utils/index.js
+++ b/demos/src/utils/index.js
@@ -1,4 +1,7 @@
/* global ImageData */
+import sum from 'lodash/sum'
+import flatten from 'lodash/flatten'
+import unpack from 'ndarray-unpack'
/**
* Find mindpoint of two points
@@ -88,9 +91,57 @@ export function centerCrop (imageData) {
}
/**
- * Takes in a ndarray of shape [x, y, z]
- * and creates ImageData layed out as [x*z, y]
+ * calculates mean and stddev for a ndarray tensor
*/
-export function flatten3DTensor (imageData) {
+export function tensorStats (tensor) {
+ const mean = sum(tensor.data) / tensor.data.length
+ const stddev = Math.sqrt(sum(tensor.data.map(x => (x - mean) ** 2)) / tensor.data.length)
+ return { mean, stddev }
+}
+/**
+ * calculates min and max for a ndarray tensor
+ */
+export function tensorMinMax (tensor) {
+ const min = Math.min(...tensor.data)
+ const max = Math.max(...tensor.data)
+ return { min, max }
+}
+
+/**
+ * Takes in a ndarray of shape [x, y]
+ * and creates image data
+ */
+export function image2Dtensor (tensor) {
+ const { min, max } = tensorMinMax(tensor)
+ let imageData = new Uint8ClampedArray(tensor.size * 4)
+ for (let i = 0, len = imageData.length; i < len; i += 4) {
+ imageData[i + 3] = 255 * (tensor.data[i / 4] - min) / (max - min)
+ }
+ return new ImageData(imageData, tensor.shape[0], tensor.shape[1])
+}
+
+/**
+ * Takes in a ndarray of shape [x, y, z]
+ * and creates an array of z ImageData [x, y] elements
+ */
+export function unroll3Dtensor (tensor) {
+ const { min, max } = tensorMinMax(tensor)
+ let shape = tensor.shape.slice()
+ let unrolled = []
+ for (let k = 0, channels = shape[2]; k < channels; k++) {
+ const channelData = flatten(unpack(tensor.pick(null, null, k)))
+ unrolled.push(channelData)
+ }
+
+ return unrolled.map(channelData => {
+ let imageData = new Uint8ClampedArray(channelData.length * 4)
+ for (let i = 0, len = channelData.length; i < len; i++) {
+ imageData[i * 4] = 0
+ imageData[i * 4 + 1] = 0
+ imageData[i * 4 + 2] = 0
+ imageData[i * 4 + 3] = 255 * (channelData[i] - min) / (max - min)
+ }
+ return new ImageData(imageData, shape[0], shape[1])
+ })
}
diff --git a/package.json b/package.json
index c6f0e99..2b8c073 100644
--- a/package.json
+++ b/package.json
@@ -34,33 +34,33 @@
},
"homepage": "https://github.com/transcranial/keras-js#readme",
"dependencies": {
- "cwise": "^1.0.9",
- "lodash": "^4.15.0",
- "ndarray": "^1.0.18",
- "ndarray-blas-level2": "^1.1.0",
- "ndarray-concat-rows": "^1.0.1",
- "ndarray-gemm": "^1.0.0",
- "ndarray-ops": "^1.2.2",
- "ndarray-squeeze": "^1.0.2",
- "ndarray-tile": "^1.0.3",
- "ndarray-unpack": "^1.0.0",
- "ndarray-unsqueeze": "^1.0.3"
+ "cwise": "1.0.9",
+ "lodash": "4.16.0",
+ "ndarray": "1.0.18",
+ "ndarray-blas-level2": "1.1.0",
+ "ndarray-concat-rows": "1.0.1",
+ "ndarray-gemm": "1.0.0",
+ "ndarray-ops": "1.2.2",
+ "ndarray-squeeze": "1.0.2",
+ "ndarray-tile": "1.0.3",
+ "ndarray-unpack": "1.0.0",
+ "ndarray-unsqueeze": "1.0.3"
},
"devDependencies": {
- "autoprefixer": "^6.4.1",
- "babel-core": "^6.14.0",
- "babel-eslint": "^6.1.2",
- "babel-loader": "^6.2.5",
- "babel-plugin-transform-class-properties": "^6.11.5",
- "babel-plugin-transform-object-rest-spread": "^6.8.0",
- "babel-polyfill": "^6.13.0",
- "babel-preset-latest": "^6.14.0",
- "css-loader": "^0.25.0",
- "http-server": "^0.9.0",
- "postcss-loader": "^0.13.0",
- "standard": "^8.1.0",
- "style-loader": "^0.13.1",
- "webpack": "^2.1.0-beta.22"
+ "autoprefixer": "6.4.1",
+ "babel-core": "6.14.0",
+ "babel-eslint": "6.1.2",
+ "babel-loader": "6.2.5",
+ "babel-plugin-transform-class-properties": "6.11.5",
+ "babel-plugin-transform-object-rest-spread": "6.8.0",
+ "babel-polyfill": "6.13.0",
+ "babel-preset-latest": "6.14.0",
+ "css-loader": "0.25.0",
+ "http-server": "0.9.0",
+ "postcss-loader": "0.13.0",
+ "standard": "8.1.0",
+ "style-loader": "0.13.1",
+ "webpack": "1.13.2"
},
"standard": {
"parser": "babel-eslint",
diff --git a/src/Model.js b/src/Model.js
index 72d61d1..5a4567d 100644
--- a/src/Model.js
+++ b/src/Model.js
@@ -274,16 +274,18 @@ export default class Model {
while (!every(inboundLayers.map(layer => layer.hasResult))) {
yield
}
+
if (layerClass === 'Merge') {
- currentLayer.result = currentLayer.call(inboundLayers.map(layer => layer.result))
- currentLayer.hasResult = true
+ currentLayer.result = currentLayer.call(inboundLayers.map(layer => {
+ return new Tensor(layer.result.tensor.data, layer.result.tensor.shape, { gpu: this.gpu })
+ }))
} else {
if (inboundLayers.length !== 1) {
throw new Error(`Layer name ${currentLayer.name} has ${inboundLayers.length} inbound nodes, but is not a Merge layer.`)
}
- currentLayer.result = currentLayer.call(inboundLayers[0].result)
- currentLayer.hasResult = true
+ currentLayer.result = currentLayer.call(new Tensor(inboundLayers[0].result.tensor.data, inboundLayers[0].result.tensor.shape, { gpu: this.gpu }))
}
+ currentLayer.hasResult = true
}
yield * this.traverseDAG(outbound)
} else {
@@ -313,7 +315,7 @@ export default class Model {
let inputLayer = this.modelLayersMap.get(inputName)
this.inputTensors[inputName] = new Tensor(inputData[inputName], inputLayer.shape, { gpu: this.gpu })
inputLayer.result = inputLayer.call(this.inputTensors[inputName])
- this.modelLayersMap.get(inputName).hasResult = true
+ inputLayer.hasResult = true
})
// start traversing DAG at input