使用tfjs-node训练MobileNet时出现model.compile is not a function错误
问题分析与修复
错误根源
你遇到的model.compile is not a function错误,核心原因是:
- 你用
tf.loadGraphModel加载的是GraphModel(计算图模型),这类模型是预训练好的推理模型,只支持预测,没有compile、fit等训练相关方法。 - 额外调用的
model.load().then()完全多余,tf.loadGraphModel本身就是异步加载完成的方法,返回的已经是加载好的模型实例。
修复方案
要微调MobileNet,需要把预训练模型作为特征提取器,在此基础上构建可训练的新模型(用LayersModel结构),具体步骤如下:
1. 正确加载预训练模型并构建微调模型
改用tf.loadLayersModel加载MobileNet的层结构模型,然后冻结预训练层(可选,减少训练参数),添加自定义的二分类输出层。
2. 移除冗余的model.load()调用
tf.loadGraphModel返回的Promise已经完成模型加载,不需要再调用model.load()。
3. 适配二分类任务
你的任务是nsfw/sfw二分类,输出层应该用units: 2配合softmax,对应sparseCategoricalCrossentropy损失函数(因为标签是整数形式)。
修复后的完整代码
import tf from '@tensorflow/tfjs-node'; import fs from 'fs'; import path from 'path'; import { promisify } from 'util'; const readdir = promisify(fs.readdir); import sharp from 'sharp'; import AWS from 'aws-sdk'; import 'dotenv/config'; AWS.config.update({ accessKeyId: process.env.AWS_ACCESS_KEY, secretAccessKey: process.env.AWS_SECRET_ACCESS_KEY, region: process.env.AWS_REGION, }); const s3 = new AWS.S3(); async function prepareDataset(datasetPath, imageDimensions) { const classNames = ['nsfw', 'sfw']; const xs = []; const ys = []; for (let i = 0; i < classNames.length; i++) { const classPath = path.join(datasetPath, classNames[i]); const files = await readdir(classPath); for (const file of files) { const imagePath = path.join(classPath, file); const imageBuffer = await sharp(imagePath) .resize(imageDimensions, imageDimensions) .toBuffer(); const image = tf.node.decodeImage(imageBuffer, 3); xs.push(image); ys.push(i); // Label (0 for 'nsfw', 1 for 'sfw') } } const xsTensor = tf.stack(xs); const ysTensor = tf.tensor1d(ys, 'int32'); const xsNormalized = xsTensor.div(255.0); const splitRatio = 0.8; // 80% for training, 20% for testing const numExamples = xsNormalized.shape[0]; const numTrainExamples = Math.floor(numExamples * splitRatio); const [xTrain, xTest] = tf.split(xsNormalized, [numTrainExamples, numExamples - numTrainExamples]); const [yTrain, yTest] = tf.split(ysTensor, [numTrainExamples, numExamples - numTrainExamples]); return { xTrain, yTrain, xTest, yTest }; } // Load your dataset from S3 to a local directory const sourceBucket = 'my-bucket'; const datasetPath = './images/'; const imageDimensions = 224; const params = { Bucket: sourceBucket, Prefix: 'dataset/' }; s3.listObjectsV2(params, async (err, data) => { if (err) { console.error('Error listing objects in S3:', err); } else { for (const object of data.Contents) { const objectKey = object.Key; const sanitizedObjectKey = objectKey.replace('dataset', ''); const localPath = path.join(datasetPath, sanitizedObjectKey); try { await fs.promises.access(localPath); console.log(`File already exists locally: ${localPath}`); } catch (error) { // Create the directory if it doesn't exist const localDirectory = path.dirname(localPath); console.log('local', localDirectory) await fs.promises.mkdir(localDirectory, { recursive: true }); const getObjectParams = { Bucket: sourceBucket, Key: objectKey }; await s3.getObject(getObjectParams).promise() .then((data) => { fs.writeFileSync(localPath, data.Body); }) .catch((error) => { console.error('Error getting object from S3:', error); }); } } prepareDataset(datasetPath, imageDimensions).then(async (data) => { // 加载MobileNet作为特征提取器(使用LayersModel结构) const baseModel = await tf.loadLayersModel('https://storage.googleapis.com/tfjs-models/tfjs/mobilenet_v3_small_075_224/model.json'); // 冻结预训练层,只训练自定义分类层(可选,加快训练速度) baseModel.trainable = false; // 构建微调模型:取MobileNet的倒数第二层作为特征输出,添加自定义分类层 const featureLayer = baseModel.getLayer('global_average_pooling2d').output; const outputLayer = tf.layers.dense({ units: 2, // 二分类,输出两个类别概率 activation: 'softmax' }).apply(featureLayer); const model = tf.model({ inputs: baseModel.inputs, outputs: outputLayer }); // 编译模型 model.compile({ optimizer: tf.train.adam(0.001), loss: 'sparseCategoricalCrossentropy', metrics: ['accuracy'], }); // 训练模型 await model.fit(data.xTrain, data.yTrain, { epochs: 10, validationData: [data.xTest, data.yTest], callbacks: { onEpochEnd: (epoch, logs) => { console.log(`Epoch ${epoch+1}: loss=${logs.loss}, accuracy=${logs.acc}, val_loss=${logs.val_loss}, val_accuracy=${logs.val_acc}`); } } }); // 保存模型到S3 await model.save(tf.io.withSaveHandler(async (modelArtifacts) => { const buffer = await tf.node.encodeModel(modelArtifacts); const targetBucket = 'my-bucket'; const targetKey = 'trained_model/model.json'; await s3.upload({ Bucket: targetBucket, Key: targetKey, Body: buffer }).promise(); console.log('Model saved to AWS S3:', targetBucket, targetKey); return { modelArtifacts }; })); }); } });
关键修改说明
- 改用
tf.loadLayersModel加载MobileNet的层结构模型,这样才能进行后续的层修改和训练。 - 冻结预训练层(
baseModel.trainable = false):如果你的数据集不大,冻结预训练层可以避免过拟合,同时大幅减少训练时间;如果数据集足够大,可以设置为true进行全量微调。 - 构建自定义输出层:适配二分类任务,使用
units:2和softmax激活函数,对应sparseCategoricalCrossentropy损失函数(因为标签是整数形式)。 - 移除了冗余的
model.load()调用,简化了异步逻辑(改用await替代链式then,代码更清晰)。
内容的提问来源于stack exchange,提问作者henrydoe
相关产品推荐
相关产品推荐

