UNPKG

tensorflow-helpers

Version:

Helper functions to use tensorflow in nodejs for transfer learning, image classification, and more

355 lines (354 loc) 12.2 kB
"use strict"; var __createBinding = (this && this.__createBinding) || (Object.create ? (function(o, m, k, k2) { if (k2 === undefined) k2 = k; var desc = Object.getOwnPropertyDescriptor(m, k); if (!desc || ("get" in desc ? !m.__esModule : desc.writable || desc.configurable)) { desc = { enumerable: true, get: function() { return m[k]; } }; } Object.defineProperty(o, k2, desc); }) : (function(o, m, k, k2) { if (k2 === undefined) k2 = k; o[k2] = m[k]; })); var __setModuleDefault = (this && this.__setModuleDefault) || (Object.create ? (function(o, v) { Object.defineProperty(o, "default", { enumerable: true, value: v }); }) : function(o, v) { o["default"] = v; }); var __importStar = (this && this.__importStar) || (function () { var ownKeys = function(o) { ownKeys = Object.getOwnPropertyNames || function (o) { var ar = []; for (var k in o) if (Object.prototype.hasOwnProperty.call(o, k)) ar[ar.length] = k; return ar; }; return ownKeys(o); }; return function (mod) { if (mod && mod.__esModule) return mod; var result = {}; if (mod != null) for (var k = ownKeys(mod), i = 0; i < k.length; i++) if (k[i] !== "default") __createBinding(result, mod, k[i]); __setModuleDefault(result, mod); return result; }; })(); Object.defineProperty(exports, "__esModule", { value: true }); exports.loadGraphModel = loadGraphModel; exports.loadLayersModel = loadLayersModel; exports.cachedLoadGraphModel = cachedLoadGraphModel; exports.cachedLoadLayersModel = cachedLoadLayersModel; exports.loadImageModel = loadImageModel; const tf = __importStar(require("@tensorflow/tfjs")); const image_utils_1 = require("../image-utils"); const tensor_1 = require("../tensor"); const classifier_utils_1 = require("../classifier-utils"); async function readFile(url) { let res = await fetch(url); let buffer = await res.arrayBuffer(); return buffer; } async function loadWeightData(file) { let buffer = await readFile(file); return new Uint8Array(buffer); } async function readJSON(url) { let res = await fetch(url); if (res.status == 404) { throw new Error('json file not found: ' + url); } let json = await res.json(); return json; } function removeModelUrlPrefix(url) { if (url.endsWith('/model.json')) { url = url.slice(0, url.length - '/model.json'.length); } if (url.endsWith('/')) { url = url.slice(0, url.length - 1); } return url; } async function getLastModified(url) { url = removeModelUrlPrefix(url); url += '/model.json'; let res = await fetch(url, { method: 'HEAD' }); if (res.status == 404) { throw new Error('file not found: ' + url); } let text = res.headers.get('Last-Modified'); return text ? new Date(text).getTime() : Date.now(); } /** * @example `loadGraphModel({ url: 'saved_model/mobilenet-v3-large-100' })` */ async function loadGraphModel(options) { let url = removeModelUrlPrefix(options.url); let classNames = options.classNames; let model = await tf.loadGraphModel({ async load() { let modelArtifact = await readJSON(url + '/model.json'); classNames = (0, classifier_utils_1.checkClassNames)(modelArtifact, classNames); let weights = modelArtifact.weightData; if (!weights) { throw new Error('missing weightData'); } if (!Array.isArray(weights)) { weights = [weights]; } for (let i = 0; i < weights.length; i++) { weights[i] = await loadWeightData(url + `/weight-${i}.bin`); } return modelArtifact; }, }); return (0, classifier_utils_1.attachClassNames)(model, classNames); } async function cachedLoadModel(args) { let { options, loadRemoteModel, loadLocalModel } = args; let { url: modelUrl, cacheUrl, classNames } = options; let localLastModified = +localStorage.getItem(cacheUrl); let checkTime = Date.now(); let remoteLastModified = options.checkForUpdates ? await getLastModified(modelUrl).catch(error => { if (localLastModified) { // skip checking if offline and already cached return localLastModified; } // throw error if offline without pre-cached copy throw error; }) : 0; if (localLastModified && (!options.checkForUpdates || localLastModified == remoteLastModified)) { try { let model = await loadLocalModel(); classNames = (0, classifier_utils_1.checkClassNames)(model, classNames); return (0, classifier_utils_1.attachClassNames)(model, classNames); } catch (error) { if (!String(error).includes('Cannot find model with path')) { throw error; } } } let _model = await loadRemoteModel(); classNames = (0, classifier_utils_1.checkClassNames)(_model, classNames); let model = (0, classifier_utils_1.attachClassNames)(_model, classNames); await model.save(cacheUrl); localStorage.setItem(cacheUrl, (remoteLastModified || checkTime).toString()); return model; } /** * @example `loadGraphModel({ url: 'saved_model/emotion-classifier' })` */ async function loadLayersModel(options) { let url = removeModelUrlPrefix(options.url); let classNames = options.classNames; let model = await tf.loadLayersModel({ async load() { let modelArtifact = await readJSON(url + '/model.json'); classNames = (0, classifier_utils_1.checkClassNames)(modelArtifact, classNames); let weights = modelArtifact.weightData; if (!weights) { throw new Error('missing weightData'); } if (!Array.isArray(weights)) { modelArtifact.weightData = await loadWeightData(url + `/weight-0.bin`); return modelArtifact; } for (let i = 0; i < weights.length; i++) { weights[i] = await loadWeightData(url + `/weight-${i}.bin`); } return modelArtifact; }, }); return (0, classifier_utils_1.attachClassNames)(model, classNames); } /** * @example ``` * cachedLoadGraphModel({ * url: 'saved_model/mobilenet-v3-large-100', * cacheUrl: 'indexeddb://mobilenet-v3-large-100', * }) * ``` */ async function cachedLoadGraphModel(options) { return cachedLoadModel({ options, loadRemoteModel: () => loadGraphModel(options), loadLocalModel: () => tf.loadGraphModel(options.cacheUrl), }); } /** * @example ``` * cachedLoadLayersModel({ * url: 'saved_model/emotion-classifier', * cacheUrl: 'indexeddb://emotion-classifier', * }) * ``` */ async function cachedLoadLayersModel(options) { return cachedLoadModel({ options, loadRemoteModel: () => loadLayersModel(options), loadLocalModel: () => tf.loadLayersModel(options.cacheUrl), }); } function getInt(str) { let int = +str; if (int && Number.isInteger(int)) { return int; } throw new TypeError(`expect int value, got: ${JSON.stringify(str)}`); } function getModelSpec(url, model) { let { signature } = model; let inputs = Object.values(signature.inputs)[0]; let outputs = Object.values(signature.outputs)[0]; let height = getInt(inputs.tensorShape.dim[1].size); let width = getInt(inputs.tensorShape.dim[2].size); let channels = getInt(inputs.tensorShape.dim[3].size); let features = getInt(outputs.tensorShape.dim[1].size); let spec = { url, width, height, channels, features, }; return spec; } function basename(url) { if (url.startsWith('data:')) { return ''; } return url.split('#')[0].split('?')[0].split('/').pop(); } function isContentHash(url) { let filename = basename(url); let ext = filename.split('.').pop(); let name = ext.length == 0 ? filename : filename.slice(0, -(ext.length + 1)); return name.length * 4 == 256 && isHexOnly(name); } function isHexOnly(str) { for (let char of str) { if (str >= '0' && str <= '9') continue; if (str >= 'A' && str <= 'F') continue; if (str >= 'a' && str <= 'f') continue; return false; } return true; } async function loadImageModel(options) { let { aspectRatio, cache } = options; let model = options.cacheUrl ? await cachedLoadGraphModel({ url: options.url, cacheUrl: options.cacheUrl, checkForUpdates: options.checkForUpdates, }) : await tf.loadGraphModel(options.url); let spec = getModelSpec(options.url, model); let { width, height, channels } = spec; async function loadImageCropped(url) { let image = new Image(); let p = new Promise((resolve, reject) => { ; (image.onload = resolve), (image.onerror = error => reject(new Error('failed to load image: ' + JSON.stringify(url), { cause: error, }))); }); image.src = url; await p; let imageTensor = tf.browser.fromPixels(image, channels); return (0, image_utils_1.cropAndResizeImageTensor)({ imageTensor, width, height, aspectRatio, }); } let fileEmbeddingCache = cache ? new Map() : null; function checkCache(url) { if (!fileEmbeddingCache || !isContentHash(url)) return; let filename = basename(url); let embedding = fileEmbeddingCache.get(filename); if (embedding) return embedding; let values = typeof cache == 'object' ? cache.get(filename) : undefined; if (!values) return; embedding = tf.tensor([values]); fileEmbeddingCache.set(filename, embedding); return embedding; } async function saveCache(file, embedding) { let filename = basename(file); fileEmbeddingCache.set(filename, embedding); if (typeof cache == 'object') { let values = Array.from(await embedding.data()); cache.set(filename, values); } } async function imageUrlToEmbedding(url) { let embedding = checkCache(url); if (embedding) return embedding; let imageTensor = await loadImageCropped(url); embedding = imageTensorToEmbedding(imageTensor); imageTensor.dispose(); if (cache && isContentHash(url)) { saveCache(url, embedding); } return embedding; } async function imageFileToEmbedding(file) { let filename = file.name; let embedding = checkCache(filename); if (embedding) return embedding; let url = await new Promise((resolve, reject) => { let reader = new FileReader(); reader.onload = () => resolve(reader.result); reader.onerror = reject; reader.readAsDataURL(file); }); let imageTensor = await loadImageCropped(url); embedding = imageTensorToEmbedding(imageTensor); imageTensor.dispose(); if (cache && isContentHash(filename)) { saveCache(filename, embedding); } return embedding; } function imageTensorToEmbedding(imageTensor) { return tf.tidy(() => { let inputTensor = (0, image_utils_1.cropAndResizeImageTensor)({ imageTensor, width, height, aspectRatio, }); let outputs = model.predict(inputTensor); let embedding = (0, tensor_1.toOneTensor)(outputs); return embedding; }); } return { spec, model, fileEmbeddingCache, checkCache, loadImageCropped, imageUrlToEmbedding, imageFileToEmbedding, imageTensorToEmbedding, }; }