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