UNPKG

tensorflow-helpers

Version:

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

427 lines (426 loc) 14.5 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 spatial_utils_1 = require("../spatial-utils"); const image_utils_1 = require("../image-utils"); const classifier_utils_1 = require("../classifier-utils"); const internal_1 = require("../internal"); const model_artifacts_1 = require("../model-artifacts"); 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).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; } async function loadModelArtifact(urlPrefix, modelArtifact) { let weightData = modelArtifact.weightData || []; if (!Array.isArray(weightData)) { weightData = [weightData]; } modelArtifact.weightData = weightData; let weightSpecs = null; if (!modelArtifact.weightSpecs) { weightSpecs = []; modelArtifact.weightSpecs = weightSpecs; } let i = 0; for (let weightsManifest of modelArtifact.weightsManifest) { for (let path of weightsManifest.paths) { let buffer = await loadWeightData(urlPrefix + '/' + path); modelArtifact.weightData[i] = buffer; i++; } if (weightSpecs) { weightSpecs.push(...weightsManifest.weights); } } } 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(); } async function getModelClassNames(url) { url = removeModelUrlPrefix(url); url += '/model.json'; let res = await fetch(url); if (res.status == 404) { throw new Error('file not found: ' + url); } let json = await res.json(); return json.classNames || undefined; } /** * @example `loadGraphModel({ url: 'saved_model/mobilenet-v3-large-100' })` */ async function loadGraphModel(options) { let url = removeModelUrlPrefix(options.url); let classNames = options.classNames; let modelArtifact = await readJSON(url + '/model.json'); (0, internal_1.patchLoadedModelJSON)(modelArtifact); classNames = (0, internal_1.checkClassNames)(modelArtifact, classNames); let model = await tf.loadGraphModel({ async load() { await loadModelArtifact(url, modelArtifact); return modelArtifact; }, }); return (0, internal_1.attachClassNames)(model, classNames); } async function cachedLoadModel(args) { let { options, loadRemoteModel, loadLocalModel } = args; let { url: modelUrl, cacheUrl, classNames } = options; let localCacheInfo = loadLocalCacheInfo(cacheUrl); let checkTime = Date.now(); let remoteLastModified = options.checkForUpdates ? await getLastModified(modelUrl).catch(error => { if (localCacheInfo.lastModified) { // skip checking if offline and already cached return localCacheInfo.lastModified; } // throw error if offline without pre-cached copy throw error; }) : 0; if (localCacheInfo.lastModified && (!options.checkForUpdates || localCacheInfo.lastModified == remoteLastModified)) { try { let model = await loadLocalModel(); model.classNames ||= localCacheInfo.classNames; classNames = (0, internal_1.checkClassNames)((0, model_artifacts_1.getModelArtifacts)(model), classNames); return (0, internal_1.attachClassNames)(model, classNames); } catch (error) { if (!String(error).includes('Cannot find model with path')) { throw error; } } } let _model = await loadRemoteModel(); _model.classNames ||= await getModelClassNames(modelUrl); classNames = (0, internal_1.checkClassNames)((0, model_artifacts_1.getModelArtifacts)(_model), classNames); let model = (0, internal_1.attachClassNames)(_model, classNames); await model.save(cacheUrl); localCacheInfo = { lastModified: remoteLastModified || checkTime, classNames, }; localStorage.setItem(cacheUrl, JSON.stringify(localCacheInfo)); return model; } function loadLocalCacheInfo(cacheUrl) { let text = localStorage.getItem(cacheUrl); let fallback = { lastModified: 0 }; try { let json = JSON.parse(text || '{}'); if (+json) { // old version, only having timestamp return fallback; } if (json.lastModified) { // new version, having timestamp and classNames return json; } return fallback; } catch (error) { return fallback; } } /** * @example `loadGraphModel({ url: 'saved_model/emotion-classifier' })` */ async function loadLayersModel(options) { let url = removeModelUrlPrefix(options.url); let classNames = options.classNames; let modelArtifact = await readJSON(url + '/model.json'); (0, internal_1.patchLoadedModelJSON)(modelArtifact); classNames = (0, internal_1.checkClassNames)(modelArtifact, classNames); let model = await tf.loadLayersModel({ async load() { await loadModelArtifact(url, modelArtifact); return modelArtifact; }, }); if (classNames) { let classCount = (0, classifier_utils_1.getClassCount)(model.outputShape); if (classCount != classNames.length) { throw new Error(`number of classes mismatch, expected: ${classNames.length}, got: ${classCount}`); } } return (0, internal_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, }; try { let name = (0, spatial_utils_1.getLastSpatialNodeName)(model); tf.tidy(() => { let input = tf.zeros([1, height, width, channels]); let output = model.execute(input, [name]); spec.spatial_features = output.shape; }); } catch (error) { // e.g. no spatial node } 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 (char >= '0' && char <= '9') continue; if (char >= 'A' && char <= 'F') continue; if (char >= 'a' && char <= 'f') continue; return false; } return true; } function loadImage(image_or_url) { if (typeof image_or_url != 'string') { return image_or_url; } let url = image_or_url; return new Promise((resolve, reject) => { let image = new Image(); image.onload = () => resolve(image); image.onerror = error => reject(new Error('failed to load image: ' + JSON.stringify(url), { cause: error, })); image.src = url; }); } async function loadImageModel(options) { let { aspectRatio, cache } = options; let model = options.cacheUrl ? await cachedLoadGraphModel({ url: options.url, cacheUrl: options.cacheUrl, checkForUpdates: options.checkForUpdates, classNames: options.classNames, }) : await loadGraphModel({ url: options.url, classNames: options.classNames, }); let spec = getModelSpec(options.url, model); let { width, height, channels } = spec; async function loadImageCropped(image_or_url) { let image = await loadImage(image_or_url); 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, options) { let embedding = checkCache(url); if (embedding) return embedding; let imageTensor = await loadImageCropped(url); embedding = imageTensorToEmbedding(imageTensor, options); imageTensor.dispose(); if (cache && isContentHash(url)) { saveCache(url, embedding); } return embedding; } async function imageFileToEmbedding(file, options) { 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, options); imageTensor.dispose(); if (cache && isContentHash(filename)) { saveCache(filename, embedding); } return embedding; } function imageTensorToEmbedding(imageTensor, options) { return tf.tidy(() => { let inputTensor = (0, image_utils_1.cropAndResizeImageTensor)({ imageTensor, width, height, aspectRatio, }); let embedding = model.predict(inputTensor); if (options?.squeeze) { embedding = tf.squeeze(embedding, [0]); } return embedding; }); } return { spec, model, fileEmbeddingCache, checkCache, loadImageCropped, imageUrlToEmbedding, imageFileToEmbedding, imageTensorToEmbedding, }; }