🧠 Stage 04 · 模型选择与搭建

选对模型,事半功倍

不要上来就堆大模型!基线模型 → 进阶模型 → 微调/蒸馏,这条路线让你稳赢90%的项目。

🌳 模型选择决策树(交互式)

告诉我你的场景,我推荐最合适的模型起点

📚 主流模型架构速查

✅ 经典图像分类任务:给整张图片打一个类别标签

ResNet

分类

残差网络,ImageNet 经典,简单有效。从 18 层到 152 层

⭐ 入门首选
torchvision.models.resnet50()

EfficientNet

分类

复合缩放,精度/效率平衡最好,B0-B7 可调

⭐ 工业首选
torchvision.models.efficientnet_b4()

Vision Transformer (ViT)

分类

把图像切成 patch 序列送入 Transformer,大数据集表现好

⭐ 新时代经典
vit_patch16_224(pretrained=True)

🔁 迁移学习:站在巨人肩膀上

绝大多数项目不需要从头训练模型!使用预训练模型 + 微调是行业标配:

全量微调

解冻所有层,全部重训

适合:数据充足(>10万)
效果最好但算力开销大
② ⭐

冻结底层+微调顶层

预训练特征提取器冻结,只训分类头

适合:数据有限(1万-10万)
性价比最高,推荐首选

特征提取 + 传统ML

用预训练模型抽特征,接XGBoost

适合:极少量数据(<1万)
工程简洁,效果稳

📝 快速迁移学习代码模板(PyTorch)

import torch
import torch.nn as nn
from torchvision import models, transforms
from torch.utils.data import DataLoader, Dataset
from PIL import Image
import torch.optim as optim

# ===== 1. 数据准备 =====
transform = transforms.Compose([
    transforms.Resize((224, 224)),
    transforms.ToTensor(),
    transforms.Normalize([0.485], [0.229])
])

class MyDataset(Dataset):
    def __init__(self, paths, labels):
        self.paths = paths; self.labels = labels
    def __getitem__(self, idx):
        img = Image.open(self.paths[idx]).convert('RGB')
        return transform(img), self.labels[idx]
    def __len__(self): return len(self.paths)

dataset = MyDataset(image_paths, labels)
loader = DataLoader(dataset, batch_size=32, shuffle=True)

# ===== 2. 加载预训练模型 =====
model = models.resnet50(pretrained=True)

# 方案B:冻结底层
for param in model.parameters():
    param.requires_grad = False

# 替换分类头
num_classes = 5
model.fc = nn.Linear(model.fc.in_features, num_classes)

# ===== 3. 训练 =====
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = model.to(device)

criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.fc.parameters(), lr=0.001)

for epoch in range(10):
    model.train()
    for images, labels in loader:
        images, labels = images.to(device), labels.to(device)
        optimizer.zero_grad()
        outputs = model(images)
        loss = criterion(outputs, labels)
        loss.backward()
        optimizer.step()
    print(ff'Epoch {epoch+1}, Loss: {loss.item():.4f}')