视频链接:https://www.bilibili.com/video/BV1aSLn6WE8f/?vd_source=5ba34935b7845cd15c65ef62c64ba82f

代码仓库:https://github.com/LitchiCheng/LLM-learning.git

大模型在任务级规划方面,扮演者世界模型的角色,对任务进行有效的分解,但大语言模型没有视觉等多模态感知的能力,类似可以替代的感知层有视觉模型,视觉语言模型 VLM,视觉生成模型(如Diffusion Model)等,今天学习一下 ViT 视觉模型。

在 ViT paper http://arxiv.org/abs/2010.11929 一文中经典的框图如下

针对左下角部分,对经典的 MINIST 数据集进行学习,单张图片为灰度图,定义 in_channels 为 1,大小 img_size 为 28

img_size = 28
patch_size = 4
in_channels = 1
embed_dim = 64
num_heads = 2
num_layers = 2
num_classes = 10
batch_size = 64
epochs = 30

class PatchEmbedding(nn.Module):
    def __init__(self, in_channels, patch_size, embed_dim):
        super().__init__()
        self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size)
    
    def forward(self, x):
        # (B, C, H, W) -> (B, embed_dim, H//patch_size, W//patch_size)
        x = self.proj(x)
        # -> (B, embed_dim, num_patches) -> (B, num_patches, embed_dim)  
        x = x.flatten(2).transpose(1, 2)
        return x

使用 Conv2d 卷积将灰度图(1通道)MINIST样本的一张图像切成若干个小块 patch,再把每个小块映射成一个向量 token(embed_dim 维特征的向量),处理完的数据 (B, num_patches, embed_dim) 下一步输入给 Transformer 作为序列数据

class MiniViT(nn.Module):
    def __init__(self):
        super().__init__()
        self.patch_embed = PatchEmbedding(in_channels, patch_size, embed_dim)
        self.cls_token = nn.Parameter(torch.randn(1, 1, embed_dim))
        self.pos_embed = nn.Parameter(torch.randn(1, num_patches + 1, embed_dim))
        
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=embed_dim, nhead=num_heads, dim_feedforward=embed_dim*2,
            batch_first=True, dropout=0.1
        )
        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)
        self.head = nn.Linear(embed_dim, num_classes)

    def forward(self, x):
        B = x.shape[0]
        x = self.patch_embed(x)
        
        cls_tokens = self.cls_token.expand(B, -1, -1)
        x = torch.cat([cls_tokens, x], dim=1)
        
        x = x + self.pos_embed
        x = self.transformer(x)
        return self.head(x[:, 0])

cls_token 用来学习分类的,后面通过 torch.cat 拼接到 patch_embed 前面,专门用来聚合全图信息,最后交给分类头判断类别

pos_embed 用来告诉 Transformer patch的位置空间信息,叠加后的 embed 就包含了空间,像素,分类特征的信息,通过反向传播学习出空间位置

encoder_layer 定义单层 Transformer

2个head,拆分掉 embed_dim,各关注一种方向的特性(笔画或者整体)

ffn 中间隐藏层维度,增加非线性表达能力

drop_out 丢掉 10% 神经元,防止过拟合

两层 encoder 用于抽象特征(鬼知道关注的啥,就是提炼出他认为抽象的内容),MINIST 比较简单,1层也可以,最后 cls_token 学习到图的所有精华抽象特征

self.head = nn.Linear(embed_dim, num_classes) 把 cls_token 学到的 64 维全局特征,映射成 10 个数字的分数,分数最高的就是预测结果

完整的训练代码,进行 MINIST 数据集的识别分类,使用 CPU 进行训练,随便一台电脑即可

import torch
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import DataLoader
from torchvision import datasets, transforms
import random
import argparse
import os

import matplotlib
matplotlib.use('TkAgg')
import matplotlib.pyplot as plt

img_size = 28
patch_size = 4
in_channels = 1
embed_dim = 64
num_heads = 2
num_layers = 2
num_classes = 10
batch_size = 64
epochs = 30

num_patches = (img_size // patch_size) ** 2


class PatchEmbedding(nn.Module):
    def __init__(self, in_channels, patch_size, embed_dim):
        super().__init__()
        self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size)
    
    def forward(self, x):
        # (B, C, H, W) -> (B, embed_dim, H//patch_size, W//patch_size)
        x = self.proj(x)
        # -> (B, embed_dim, num_patches) -> (B, num_patches, embed_dim)  
        x = x.flatten(2).transpose(1, 2)
        return x


class MiniViT(nn.Module):
    def __init__(self):
        super().__init__()
        self.patch_embed = PatchEmbedding(in_channels, patch_size, embed_dim)
        self.cls_token = nn.Parameter(torch.randn(1, 1, embed_dim))
        self.pos_embed = nn.Parameter(torch.randn(1, num_patches + 1, embed_dim))
        
        encoder_layer = nn.TransformerEncoderLayer(
            d_model=embed_dim, nhead=num_heads, dim_feedforward=embed_dim*2,
            batch_first=True, dropout=0.1
        )
        self.transformer = nn.TransformerEncoder(encoder_layer, num_layers=num_layers)
        self.head = nn.Linear(embed_dim, num_classes)

    def forward(self, x):
        B = x.shape[0]
        x = self.patch_embed(x)
        
        cls_tokens = self.cls_token.expand(B, -1, -1)
        x = torch.cat([cls_tokens, x], dim=1)
        
        x = x + self.pos_embed
        x = self.transformer(x)
        return self.head(x[:, 0])


def get_data_loaders():
    transform = transforms.Compose([
        transforms.Resize((28,28)),
        transforms.ToTensor()
    ])

    train_dataset = datasets.MNIST(root='./.data', train=True, download=True, transform=transform)
    test_dataset = datasets.MNIST(root='./.data', train=False, download=True, transform=transform)

    small_train = torch.utils.data.Subset(train_dataset, range(0, 3000))
    train_loader = DataLoader(small_train, batch_size=batch_size, shuffle=True)
    test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)
    
    return train_loader, test_loader, test_dataset


def train(train_loader, model_path, device):
    model = MiniViT().to(device)

    criterion = nn.CrossEntropyLoss()
    optimizer = optim.Adam(model.parameters(), lr=1e-3)

    print("开始训练...")
    for epoch in range(epochs):
        model.train()
        total_loss = 0
        for img, label in train_loader:
            img, label = img.to(device), label.to(device)
            out = model(img)
            loss = criterion(out, label)
            
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            total_loss += loss.item()
        
        print(f"Epoch {epoch+1}, Loss: {total_loss/len(train_loader):.4f}")

    torch.save(model.state_dict(), model_path)
    print(f"模型已保存为 {model_path}")


def eval_acc(model, loader, device):
    model.eval()
    correct = 0
    total = 0
    with torch.no_grad():
        for img, label in loader:
            img, label = img.to(device), label.to(device)
            out = model(img)
            pred = torch.argmax(out, dim=1)
            correct += (pred == label).sum().item()
            total += label.size(0)
    return correct / total


def load_model(model_path, device):
    model = MiniViT().to(device)
    if os.path.exists(model_path):
        model.load_state_dict(torch.load(model_path, map_location=device))
    else:
        print(f"警告: 模型文件 {model_path} 不存在,将使用随机初始化的权重")
    return model


def predict_and_show(model, test_dataset, device):
    idx = random.randint(0, len(test_dataset)-1)
    img, true_label = test_dataset[idx]
    
    model.eval()
    with torch.no_grad():
        img_input = img.unsqueeze(0).to(device)
        out = model(img_input)
        pred_label = torch.argmax(out, dim=1).item()
    
    plt.figure(figsize=(3,3))
    plt.imshow(img.squeeze(), cmap="gray")
    
    title = f"True: {true_label} | Pred: {pred_label}"
    plt.title(title, color="green" if true_label==pred_label else "red")
    
    plt.axis("off")
    plt.show()


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument('--mode', type=str, default='train', choices=['train', 'test'],
                        help='运行模式: train 或 test')
    parser.add_argument('--model_path', type=str, default='minivit_mnist.pth',
                        help='模型保存/加载路径')
    args = parser.parse_args()

    device = torch.device("cpu")
    train_loader, test_loader, test_dataset = get_data_loaders()

    if args.mode == 'train':
        train(train_loader, args.model_path, device)
        model = load_model(args.model_path, device)
        acc = eval_acc(model, test_loader, device)
        print(f"测试集准确率: {acc:.4f}")
        predict_and_show(model, test_dataset, device)
    else:
        model = load_model(args.model_path, device)
        acc = eval_acc(model, test_loader, device)
        print(f"测试集准确率: {acc:.4f}")
        predict_and_show(model, test_dataset, device)


if __name__ == '__main__':
    main()
python ViT.py --mode train --model_path .models/minivit.pth

python ViT.py --mode test --model_path .models/minivit.pth

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐