基于 transformer 框架而实现的 segmenter 模型

发布于 2026-06-13 21:00 1119 字 6 min read

一个跟着教程从零搭建的图像分割小模型,基于 transformer 框架整合 ViT 编码器与 MaskTransformer 解码器,使用 PASCAL VOC 2012 数据集训练。

一个自己搭建的小模型罢了,不足为奇。

前言及项目介绍

这个项目也是足足写了将近半个学期,刚开始只是跟着教程学习了一下 transformer 的各个部件,自注意力机制、embedding、位置编码这些啥的,但是一直不能把他们拼成一块,趁这个大好的机会,就把他们整合起来形成了一个谁都会写的小模型。

即使很入门,但是还是很满意的。

个人仓库地址:cateude/transformer_segment

个人认为的优点:

  1. 文档习惯较好(不过是创建了一个 claude.md)
  2. 运用 PASCAL VOC 2012,数据来源严谨,明确的 train 和 val 进行分离(没有在测试集里面参杂训练集的参数噢,这是不对的噢)
  3. 接口规范,参数可控,可记录,比硬编码的脚本更容易追溯
  4. 设备处理比较规范,嗯,就这样

但是除此以外,缺点还是比较多:

  1. 验证只跑前 20 个 batch(对不起,我的电脑做不到)
  2. 数据管道默认了理想输入,写的时候真的没有考虑到(PS:我错了)
  3. 配置 config 是所有的参数,但是其他的代码块里面的 class 定义也有参数,参数调起来相对较麻烦

项目的框架

项目框架
项目框架

项目正文的介绍

模型的分支搭建:vit 和 decoder

首先先搭建了 vit 和 decoder 两个文件,一个作为编码,一个作为解码。

def __init__(self, num_patches: int = 576, embed_dim: int = 384, depth: int = 8,
             num_heads: int = 6, patch_size: int = 16, pretrained: bool = True):
    super().__init__()
    if embed_dim not in _TIMM_NAME:
        raise ValueError(f"没有对应 embed_dim={embed_dim} 的 timm 模型,支持 {list(_TIMM_NAME)}")
    image_size = int(num_patches ** 0.5) * patch_size

    self.backbone = timm.create_model(
        _TIMM_NAME[embed_dim],
        pretrained=pretrained,
        num_classes=0,
        img_size=image_size,
    )
    self.num_prefix_tokens = self.backbone.num_prefix_tokens

def forward(self, x):
    tokens = self.backbone.forward_features(x)
    return tokens[:, self.num_prefix_tokens:, :]

PS:由于训练效果的原因,不得已加入了 timm 模块进行辅助,呜呜!

模型的组合:model

在 model.py 里面把这两个进行了拼接,并且加了可学习温度参数。

class Segmenter(nn.Module):
    def __init__(self, num_classes = 21, image_size = 384, patch_size = 16,
                 embed_dim = 768, encoder_depth = 12, encoder_heads = 12,
                 decoder_depth = 4, decoder_heads = 12):
        super().__init__()
        self.num_classes = num_classes
        self.image_size = image_size
        self.patch_size = patch_size
        num_patches = (image_size // patch_size) ** 2
        self.encoder = VIT(num_patches = num_patches, embed_dim = embed_dim, depth = encoder_depth, num_heads = encoder_heads, patch_size = patch_size)
        self.decoder = MaskTransformer(embed_dim = embed_dim, num_heads = decoder_heads, num_classes = num_classes, depth = decoder_depth)
        self.logit_scale = nn.Parameter(torch.log(torch.tensor(1.0 / 0.07)))

    def forward(self, x):
        B = x.shape[0]
        patch_tokens = self.encoder(x)
        cls_embedding = self.decoder(patch_tokens)

        patch_tokens = F.normalize(patch_tokens, dim = -1)
        cls_embedding = F.normalize(cls_embedding, dim = -1)
        logit_scale = self.logit_scale.exp().clamp(max = 100.0)
        logits = logit_scale * (patch_tokens @ cls_embedding.transpose(1, 2))
        logits = logits.permute(0, 2, 1)

        H_patch = W_patch = x.shape[2] // self.patch_size
        logits = logits.reshape(B, self.num_classes, H_patch, W_patch)

        logits = F.interpolate(logits, size = (self.image_size, self.image_size), mode = 'bilinear', align_corners = False)

        return logits

数据的处理:dataset 和 transforms

dataset 分离训练集和测试集,且加入了图片的读取。

transforms 里面更是老三样(Resize、ToTensor、Normalize,不过在训练集的处理上加了 RandomCrop 和 RandomHorizontalFlip)。

class Resize:
    def __init__(self, size):
        self.size = size

    def __call__(self, sample):
        image = sample["image"]
        mask = sample["mask"]
        image = F.resize(image, self.size, InterpolationMode.BILINEAR)
        mask = F.resize(mask, self.size, InterpolationMode.NEAREST) if mask is not None else None
        return {"image": image, "mask": mask}

class ToTensor:
    def __call__(self, sample):
        image = sample["image"]
        mask = sample["mask"]
        image = F.to_tensor(image)
        mask = torch.as_tensor(np.array(mask), dtype=torch.long) if mask is not None else None
        return {"image": image, "mask": mask}

class Normalize:
    def __init__(self, mean, std):
        self.mean = mean
        self.std = std

    def __call__(self, sample):
        image = sample["image"]
        mask = sample["mask"]
        image = F.normalize(image, mean=self.mean, std=self.std)
        return {"image": image, "mask": mask}

损失函数和评估:losses 和 metrics

不要搞混哟,losses 是给模型看的,是训练的时候进行评估的。metrics 是给屏幕前的你看的,是验证测试的时候评估的。

单步测试:train_step

一个单步训练,放着前向、反向、优化器更新,经典的训练五部曲。

def train_step(model, images, masks, optimizer, scaler, criterion) -> float:
    optimizer.zero_grad()
    device_type = images.device.type

    with torch.amp.autocast(device_type):
        logits = model(images)
        loss = criterion(logits, masks)

    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()

    return loss.item()

测试:train

把所有的模块整合到一起进行训练模型。

结果呈现:predict

呈现最后的结果。

预测结果
预测结果

如何启动

很简单:

  1. 下载 uv 包 pip install uv 或者其他更快的方式
  2. 安装匹配你的 gpu 包(PS:被这个硬控了一天,那个时候真是绝望啊)
  3. 终端输入 uv run python train.py
  4. 静候等待其训练完成(大概仅需 1 个小时)
  5. 终端输入 uv run python predict.py,然后你的 output 里会生成出来一张图

总结

作为一个刚开始接触深度学习的小白,学会了从看教程到复现搭模型的跨越 (⊙﹏⊙)(Ciallo~(∠・ω< )⌒★