tensorflow-helpers
Version:
Helper functions to use tensorflow in nodejs for transfer learning, image classification, and more
355 lines (354 loc) • 12.2 kB
JavaScript
;
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,
};
}