diff --git a/demos/src/imdb-bidirectional-lstm.css b/demos/src/imdb-bidirectional-lstm.css new file mode 100644 index 0000000..e2e7fe0 --- /dev/null +++ b/demos/src/imdb-bidirectional-lstm.css @@ -0,0 +1,133 @@ +@import './_variables.css'; + +.demo.imdb-bidirectional-lstm { + .column { + display: flex; + align-items: center; + justify-content: center; + } + + .column.input-column { + justify-content: center; + + .input-container { + text-align: right; + margin: 5px 5px 5px 20px; + position: relative; + + .input-label { + font-family: $font-3; + font-size: 18px; + color: $color-2; + text-align: left; + } + + .mdl-textfield { + width: 550px; + + textarea { + color: $color-3; + font-family: $font-1; + font-size: 18px; + padding: 10px; + } + } + + .input-clear { + display: flex; + align-items: center; + justify-content: flex-end; + color: $color-2; + transition: color 0.2s ease-in; + + &:hover { + color: $color-1-lighter; + cursor: pointer; + } + } + } + } + + .column.output-column { + justify-content: center; + + .output { + height: 160px; + display: flex; + flex-direction: row; + align-items: flex-end; + justify-content: center; + user-select: none; + cursor: default; + + .output-class { + display: flex; + flex-direction: column; + align-items: center; + justify-content: center; + padding: 0 6px; + border-bottom: 2px solid $color-1-lighter; + + .output-label { + font-family: $font-2; + font-size: 1.5rem; + color: $color-2; + } + + .output-bar { + width: 8px; + background: #EEEEEE; + transition: height 0.2s ease-out; + } + } + + .output-class.predicted { + border-bottom-color: $color-1; + + .output-label { + color: $color-1; + } + } + } + } + + .layer-results-container { + position: relative; + + .layer-result { + position: relative; + z-index: 1; + margin: 30px 20px; + background: white; + border-radius: 10px; + padding: 20px; + overflow-x: auto; + + .layer-result-heading { + font-size: 1rem; + color: #999999; + margin-bottom: 10px; + display: flex; + flex-direction: column; + font-size: 12px; + + span.layer-class { + color: $color-1; + font-size: 14px; + font-weight: bold; + } + } + + .layer-result-canvas-container { + display: inline-flex; + flex-wrap: wrap; + background: white; + + canvas { + border: 1px solid lightgray; + margin: 1px; + } + } + } + } +} diff --git a/demos/src/imdb-bidirectional-lstm.js b/demos/src/imdb-bidirectional-lstm.js new file mode 100644 index 0000000..04ce8d3 --- /dev/null +++ b/demos/src/imdb-bidirectional-lstm.js @@ -0,0 +1,215 @@ +/* global Vue */ +import './imdb-bidirectional-lstm.css' + +import debounce from 'lodash/debounce' +import * as utils from './utils' + +const MODEL_FILEPATHS_DEV = { + model: '/demos/data/imdb_bidirectional_lstm/imdb_bidirectional_lstm.json', + weights: '/demos/data/imdb_bidirectional_lstm/imdb_bidirectional_lstm_weights.buf', + metadata: '/demos/data/imdb_bidirectional_lstm/imdb_bidirectional_lstm_metadata.json' +} +const MODEL_FILEPATHS_PROD = { + model: 'demos/data/imdb_bidirectional_lstm/imdb_bidirectional_lstm.json', + weights: 'https://transcranial.github.io/keras-js-demos-data/imdb_bidirectional_lstm/imdb_bidirectional_lstm_weights.buf', + metadata: 'demos/data/imdb_bidirectional_lstm/imdb_bidirectional_lstm_metadata.json' +} +const MODEL_CONFIG = { + filepaths: (process.env.NODE_ENV === 'production') ? MODEL_FILEPATHS_PROD : MODEL_FILEPATHS_DEV +} + +const LAYER_DISPLAY_CONFIG = { +} + +/** + * + * VUE COMPONENT + * + */ +export const ImdbBidirectionalLstm = Vue.extend({ + props: ['hasWebgl'], + + template: require('raw!./imdb-bidirectional-lstm.template.html'), + + data: function () { + return { + model: new KerasJS.Model(Object.assign({ gpu: this.hasWebgl }, MODEL_CONFIG)), + modelLoading: true, + input: new Float32Array(200), + output: new Float32Array(1), + layerResultImages: [], + layerDisplayConfig: LAYER_DISPLAY_CONFIG, + drawing: false, + strokes: [], + useGpu: this.hasWebgl + } + }, + + computed: { + loadingProgress: function () { + return this.model.getLoadingProgress() + } + }, + + ready: function () { + this.model.ready().then(() => { + this.modelLoading = false + this.$nextTick(function () { + this.getIntermediateResults() + }) + }) + }, + + methods: { + + toggleGpu: function () { + this.model.toggleGpu(!this.useGpu) + }, + + 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') + ctxCenterCrop.clearRect(0, 0, ctxCenterCrop.canvas.width, ctxCenterCrop.canvas.height) + const ctxScaled = document.getElementById('input-canvas-scaled').getContext('2d') + ctxScaled.clearRect(0, 0, ctxScaled.canvas.width, ctxScaled.canvas.height) + this.output = new Float32Array(10) + this.drawing = false + this.strokes = [] + }, + + activateDraw: function (e) { + this.drawing = true + this.strokes.push([]) + let points = this.strokes[this.strokes.length - 1] + points.push(utils.getCoordinates(e)) + }, + + draw: function (e) { + if (!this.drawing) return + + const ctx = document.getElementById('input-canvas').getContext('2d') + + ctx.lineWidth = 20 + ctx.lineJoin = ctx.lineCap = 'round' + ctx.strokeStyle = '#393E46' + + ctx.clearRect(0, 0, ctx.canvas.width, ctx.canvas.height) + + let points = this.strokes[this.strokes.length - 1] + points.push(utils.getCoordinates(e)) + + // draw individual strokes + for (let s = 0, slen = this.strokes.length; s < slen; s++) { + points = this.strokes[s] + + let p1 = points[0] + let p2 = points[1] + ctx.beginPath() + ctx.moveTo(...p1) + + // draw points in stroke + // quadratic bezier curve + for (let i = 1, len = points.length; i < len; i++) { + ctx.quadraticCurveTo(...p1, ...utils.getMidpoint(p1, p2)) + p1 = points[i] + p2 = points[i + 1] + } + ctx.lineTo(...p1) + ctx.stroke() + } + }, + + deactivateDrawAndPredict: debounce(function () { + if (!this.drawing) return + this.drawing = false + + const ctx = document.getElementById('input-canvas').getContext('2d') + + // center crop + const imageDataCenterCrop = utils.centerCrop(ctx.getImageData(0, 0, ctx.canvas.width, ctx.canvas.height)) + const ctxCenterCrop = document.getElementById('input-canvas-centercrop').getContext('2d') + ctxCenterCrop.canvas.width = imageDataCenterCrop.width + ctxCenterCrop.canvas.height = imageDataCenterCrop.height + ctxCenterCrop.putImageData(imageDataCenterCrop, 0, 0) + + // scaled to 28 x 28 + const ctxScaled = document.getElementById('input-canvas-scaled').getContext('2d') + ctxScaled.save() + ctxScaled.scale(28 / ctxCenterCrop.canvas.width, 28 / ctxCenterCrop.canvas.height) + ctxScaled.clearRect(0, 0, ctxCenterCrop.canvas.width, ctxCenterCrop.canvas.height) + ctxScaled.drawImage(document.getElementById('input-canvas-centercrop'), 0, 0) + const imageDataScaled = ctxScaled.getImageData(0, 0, ctxScaled.canvas.width, ctxScaled.canvas.height) + ctxScaled.restore() + + // process image data for model input + const { data } = imageDataScaled + this.input = new Float32Array(784) + for (let i = 0, len = data.length; i < len; i += 4) { + this.input[i / 4] = data[i + 3] / 255 + } + + this.model.predict({ input: this.input }).then(outputData => { + this.output = outputData.output + this.getIntermediateResults() + }) + }, 200, { leading: true, trailing: true }), + + getIntermediateResults: function () { + let results = [] + for (let [name, layer] of this.model.modelLayersMap.entries()) { + if (name === 'input') continue + + const layerClass = layer.layerClass || '' + + let images = [] + if (layer.result && layer.result.tensor.shape.length === 3) { + images = utils.unroll3Dtensor(layer.result.tensor) + } else if (layer.result && layer.result.tensor.shape.length === 2) { + images = [utils.image2Dtensor(layer.result.tensor)] + } else if (layer.result && layer.result.tensor.shape.length === 1) { + images = [utils.image1Dtensor(layer.result.tensor)] + } + results.push({ + name, + layerClass, + images + }) + } + this.layerResultImages = results + setTimeout(() => { + this.showIntermediateResults() + }, 0) + }, + + showIntermediateResults: function () { + this.layerResultImages.forEach((result, layerNum) => { + const scalingFactor = this.layerDisplayConfig[result.name].scalingFactor + result.images.forEach((image, imageNum) => { + const ctx = document.getElementById(`intermediate-result-${layerNum}-${imageNum}`).getContext('2d') + ctx.putImageData(image, 0, 0) + const ctxScaled = document.getElementById(`intermediate-result-${layerNum}-${imageNum}-scaled`).getContext('2d') + ctxScaled.save() + ctxScaled.scale(scalingFactor, scalingFactor) + 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) => { + const scalingFactor = this.layerDisplayConfig[result.name].scalingFactor + result.images.forEach((image, imageNum) => { + const ctxScaled = document.getElementById(`intermediate-result-${layerNum}-${imageNum}-scaled`).getContext('2d') + ctxScaled.save() + ctxScaled.scale(scalingFactor, scalingFactor) + ctxScaled.clearRect(0, 0, ctxScaled.canvas.width, ctxScaled.canvas.height) + ctxScaled.restore() + }) + }) + } + } +}) diff --git a/demos/src/imdb-bidirectional-lstm.template.html b/demos/src/imdb-bidirectional-lstm.template.html new file mode 100644 index 0000000..c8da85e --- /dev/null +++ b/demos/src/imdb-bidirectional-lstm.template.html @@ -0,0 +1,63 @@ +