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

如何访问与使用CNTK迁移学习生成的TransferLearning.model文件?

如何访问和使用CNTK迁移学习生成的TransferLearning.model

一、加载模型进行预测

训练好的.model文件是CNTK的序列化模型,你可以用CNTK的Python API直接加载它,然后对新图像做预测。这里给你一个实用的代码示例,记得和你训练时的预处理逻辑保持一致:

import cntk as C
import numpy as np
from PIL import Image

# 加载训练好的模型,替换成你的实际路径
model = C.load_model("~/CNTK-Samples-2-3-1/Examples/Image/TransferLearning/Output/TransferLearning.model")

# 图像预处理函数,必须和训练时的处理完全匹配
def preprocess_image(image_path):
    # 调整图像大小为训练时的输入尺寸(教程里一般是224x224)
    img = Image.open(image_path).resize((224, 224))
    # 转成numpy数组并归一化
    img_data = np.array(img, dtype=np.float32) / 255.0
    # CNTK默认使用通道优先(CHW)格式,所以转置维度
    img_data = np.transpose(img_data, (2, 0, 1))
    # 增加batch维度(模型接受批量输入)
    img_data = np.expand_dims(img_data, axis=0)
    return img_data

# 测试单张图像
test_image_path = "你的测试图片路径"
input_tensor = preprocess_image(test_image_path)
# 执行预测
pred_results = model.eval(input_tensor)

# 解析结果,替换成你自定义数据集的类别名称
class_labels = ["类别A", "类别B", "类别C"]  # 改成你自己的类别
predicted_idx = np.argmax(pred_results)
print(f"预测类别: {class_labels[predicted_idx]}, 置信度: {pred_results[0][predicted_idx]:.4f}")

如果你不想写代码,也可以直接修改教程里的预测脚本,把模型路径指向你生成的TransferLearning.model,直接运行脚本就能得到预测结果。

二、修改或继续训练模型

如果想要调整模型结构(比如更换分类头)或者继续微调模型,可以这样操作:

import cntk as C

# 加载现有模型
model = C.load_model("~/CNTK-Samples-2-3-1/Examples/Image/TransferLearning/Output/TransferLearning.model")

# 获取模型的特征提取部分(假设最后一层是分类层,我们取倒数第二层的输出作为特征)
feature_extractor = model.layers[-2].output

# 构建新的分类层(比如调整类别数)
new_num_classes = 3
new_classification_layer = C.layers.Dense(new_num_classes, activation=C.softmax)(feature_extractor)

# 组合成新的模型
updated_model = C.combine(new_classification_layer)

# 之后就可以用新的数据集继续训练这个模型,设置损失函数、优化器等
# 比如定义损失和优化器
input_var = model.inputs[0]
label_var = C.input_variable(new_num_classes)
loss = C.cross_entropy_with_softmax(updated_model(input_var), label_var)
learner = C.adam(updated_model.parameters, lr=C.learning_rate_schedule(0.001, C.UnitType.minibatch))
trainer = C.Trainer(updated_model, (loss, C.classification_error(updated_model(input_var), label_var)), [learner])

三、注意事项

  • 预处理逻辑必须和训练时完全一致(图像大小、归一化方式、通道顺序),否则预测结果会失真。
  • 如果不确定模型的输入输出结构,可以用model.inputs和model.outputs查看模型的输入输出变量,也可以用model.summary()打印模型结构。

内容的提问来源于stack exchange,提问作者BBS

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 08:08:01