From 4b04c0496dba6b14ae709e52b7b3f73e78eae793 Mon Sep 17 00:00:00 2001 From: Leon Chen Date: Tue, 20 Sep 2016 01:31:21 -0400 Subject: [PATCH] update mnist convnet demo --- demos/src/mnist-cnn.css | 40 +++++++++++++++++++++ demos/src/mnist-cnn.js | 76 ++++++++++++++++++++++++++++++++++++++-- demos/src/utils/index.js | 57 ++++++++++++++++++++++++++++-- package.json | 50 +++++++++++++------------- src/Model.js | 12 ++++--- 5 files changed, 199 insertions(+), 36 deletions(-) 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