tensorflow-helpers
Version:
Helper functions to use tensorflow in nodejs for transfer learning, image classification, and more
117 lines (111 loc) • 4.24 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 });
const tf = __importStar(require("@tensorflow/tfjs"));
const fs_1 = require("fs");
const path_1 = require("path");
const model_1 = require("./model");
const model_artifacts_1 = require("./model-artifacts");
let helpMessage = `
Usage: download-tfjs-model <source> <output-dir>
Download and save TensorFlow.js models for use in browser or Node.js.
Supports URLs (TensorFlow Hub, Kaggle) and local model files/directories.
Examples:
# Download MobileNet V3
download-tfjs-model https://www.kaggle.com/models/google/mobilenet-v3/TfJs/large-100-224-feature-vector/1 ./browser-models/mobilenet-v3-large-100
# Download MobileNet V2
download-tfjs-model https://www.kaggle.com/models/google/mobilenet-v2/TfJs/035-128-feature-vector/3 ./browser-models/mobilenet-v2-035
# Convert local model
download-tfjs-model ./hub-models/mobilenet-v2-035-128-feature-vector ./browser-models/mobilenet-v2-035
`.trim();
async function main() {
const args = process.argv.slice(2);
if (args.includes('--help') || args.includes('-h')) {
console.log(helpMessage);
process.exit(0);
}
if (args.length != 2) {
console.error(helpMessage);
process.exit(1);
}
const [source, outputDir] = args;
try {
console.log(`Loading model from: ${source}`);
const model = await loadModel(source);
let artifacts = (0, model_artifacts_1.getModelArtifacts)(model);
artifacts.userDefinedMetadata ||= {};
artifacts.userDefinedMetadata.source = source;
console.log(`Saving model to: ${outputDir}`);
(0, fs_1.mkdirSync)(outputDir, { recursive: true });
await (0, model_1.saveModel)({ model, dir: outputDir });
console.log('✅ Model downloaded and saved successfully!');
}
catch (error) {
console.error('❌ Error:', error);
process.exit(1);
}
}
async function loadModel(source) {
if (!source.startsWith('http') && !source.endsWith('.json')) {
source = (0, path_1.join)(source, 'model.json');
}
if (!source.includes('://')) {
source = `file://${source}`;
}
try {
// try layered model
if (source.startsWith('file://')) {
return await (0, model_1.loadLayersModel)({ dir: toDir(source) });
}
else {
return await tf.loadLayersModel(source, { fromTFHub: true });
}
}
catch {
// try graph model
if (source.startsWith('file://')) {
return await (0, model_1.loadGraphModel)({ dir: toDir(source) });
}
else {
return await tf.loadGraphModel(source, { fromTFHub: true });
}
}
}
function toDir(source) {
source = source.replace('file://', '');
return (0, path_1.dirname)(source);
}
main();