UNPKG

tensorflow-helpers

Version:

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

65 lines (64 loc) 2.44 kB
"use strict"; Object.defineProperty(exports, "__esModule", { value: true }); exports.patchLoadedModelJSON = patchLoadedModelJSON; exports.checkClassNames = checkClassNames; exports.attachClassNames = attachClassNames; const model_artifacts_1 = require("./model-artifacts"); function patchLoadedModelJSON(model) { let changed = false; // recover weightsManifest from inline weightData and weightSpecs if (model.weightData && model.weightSpecs && !model.weightsManifest) { let paths = []; let weightData = model.weightData; if (!Array.isArray(weightData)) { weightData = [weightData]; } let n = weightData.length; for (let i = 0; i < n; i++) { let filename = `weight-${i}.bin`; paths.push(filename); } model.weightsManifest = [{ paths, weights: model.weightSpecs }]; delete model.weightData; delete model.weightSpecs; changed = true; } // move the classNames to userDefinedMetadata if (model.classNames) { model.userDefinedMetadata ||= {}; model.userDefinedMetadata.classNames = model.classNames; delete model.classNames; changed = true; } return changed; } function checkClassNames(modelArtifact, classNames) { let classNamesInMetadata = modelArtifact.userDefinedMetadata?.classNames; if (classNamesInMetadata) { if (!Array.isArray(classNamesInMetadata)) { throw new Error('classNames in userDefinedMetadata is not an array'); } if (typeof classNamesInMetadata[0] !== 'string') { throw new Error('classNames in userDefinedMetadata is not an array of strings'); } } if (classNames && classNamesInMetadata) { let expected = JSON.stringify(classNames); let actual = JSON.stringify(classNamesInMetadata); if (actual !== expected) { throw new Error(`classNames mismatch, expected: ${expected}, actual: ${actual}`); } } return !classNames && classNamesInMetadata ? classNamesInMetadata : classNames; } function attachClassNames(_model, classNames) { let model = (0, model_artifacts_1.exposeModelArtifacts)(_model); if (classNames) { let artifacts = model.getArtifacts(); artifacts.userDefinedMetadata ||= {}; artifacts.userDefinedMetadata.classNames = classNames; } return model; }