import torch
import torchvision
from torch import nn

# ==================== 1. 加载预训练模型 ====================
print("=" * 50)
print("1. 加载预训练 VGG16 模型")
vgg16 = torchvision.models.vgg16(pretrained=True)
print(vgg16)  # 会打印出完整的网络结构

# ==================== 2. 修改模型 (适配 CIFAR-10) ====================
print("\n" + "=" * 50)
print("2. 修改最后一层为 10 分类 (适配 CIFAR-10)")

# 方法1:直接替换(推荐)
vgg16.classifier[6] = nn.Linear(4096, 10)

# 方法2:追加一层
# vgg16.classifier.add_module('my_linear', nn.Linear(1000, 10))

print(vgg16)  # 打印修改后的结构

# ==================== 3. 模型保存 (P26) ====================
print("\n" + "=" * 50)
print("3. 保存模型")

# 方式1:保存完整模型(结构 + 参数)
torch.save(vgg16, "vgg16_method1.pth")

# 方式2:只保存参数(推荐)
torch.save(vgg16.state_dict(), "vgg16_method2.pth")

print("模型已保存!")

# ==================== 4. 模型加载 (P26) ====================
print("\n" + "=" * 50)
print("4. 加载模型")

# 加载方式1:加载完整模型
# 注意:这样加载需要模型类的定义在当前文件中(如果是自定义模型,需要先定义类)
model1 = torch.load("vgg16_method1.pth")
print("方式1加载成功!")

# 加载方式2:先重建模型结构,再加载参数(推荐)
# 注意:重建时必须与保存时的结构完全一致!
vgg16_reload = torchvision.models.vgg16(pretrained=False)
vgg16_reload.classifier[6] = nn.Linear(4096, 10)  # 先做同样的修改
vgg16_reload.load_state_dict(torch.load("vgg16_method2.pth"))
print("方式2加载成功!")

# ==================== 5. 验证参数一致 ====================
print("\n" + "=" * 50)
print("5. 验证参数是否一致")

# 比较两种方式加载的模型的第一层卷积核是否相同
same = torch.equal(model1.features[0].weight, vgg16_reload.features[0].weight)
print(f"两种方式加载的模型参数是否一致? {same}")

print("\n✅ 所有操作完成!")

1.vgg16 = torchvision.models.vgg16(pretrained=True)

pretrained=True        模型vgg16已经拿超级多的图片训练过了=一个见多识广的老人

pretrained=Flase        模型vgg16没训练过=一个小孩子

2.原来是Linear(in_features=4096, out_features=1000, bias=True)

修改有两种方式,方式一直接修改:vgg16.classifier[6] = nn.Linear(4096, 10)

Linear(in_features=4096, out_features=10, bias=True)

方式二vgg16.classifier.add_module('my_linear', nn.Linear(1000, 10))

这是打补丁,推荐使用方式一。

3.保存有两种方式,方式一保存结构和参数:torch.save(vgg16, "vgg16_method1.pth")

方式二用字典只保存参数:torch.save(vgg16.state_dict(), "vgg16_method2.pth")

相当于你要去图书馆看书,方式一是你把连书带书架全都借走了,方式二是只借书,推荐方式二

4.但是使用方式二,你再次加载时需要重新加载骨架,因为你保存时就没保存骨架,像这样:

vgg16_reload = torchvision.models.vgg16(pretrained=False) vgg16_reload.classifier[6] = nn.Linear(4096, 10) # 先做同样的修改 vgg16_reload.load_state_dict(torch.load("vgg16_method2.pth"))

两种保存实质保存的是一个内容,所以卷积核一样

5.为什么要保存pth文件,直接用python文件不好吗?

python文件保存的是目录,不是内容,不是具体的卷积核参数,你保存python文件下一次还得重新加载。保存pth模型文件,记录了上次训练的结果,下次拿来可以直接用

Logo

中国智能体开发者社区,聚焦智能体与大模型开发,提供前沿资讯、实用工具链、开源项目及行业案例。通过技术沙龙、开发者大赛等活动,促进经验交流与协作,助力开发者快速构建创新智能应用。

更多推荐