You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

使用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.07 09:26:02