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

如何在运行时获取ResNet模型全连接层输出用于UMAP可视化

问题解答

是否可以通过在__init__中添加空列表存储全连接层输出

可以,这种方案完全可行,只需要注意两个细节避免出问题:

  • 直接存储带计算图的x_hat会占用大量显存,存储前需要用detach()剥离计算图,再转存到CPU内存,避免显存溢出
  • 多卡训练/测试时,直接存在模型实例列表里的只有当前卡的输出,需要用Lightning自带的self.all_gather聚合所有卡的输出,才能拿到完整的全量数据

可直接复用的修改后代码示例

import torch
import torch.nn as nn
import torchvision
import pytorch_lightning as pl

# 原模型定义不变
def get_model():
    model = torchvision.models.resnet50(pretrained=True)
    num_ftrs = model.fc.in_features
    model.fc = nn.Linear(num_ftrs, 2)
    return model

class classifierModel(pl.LightningModule):
    def __init__(self, model):
        super().__init__()
        self.model = model
        self.learning_rate = 0.0001
        # 初始化空列表分别存储输出和对应标签,方便后续UMAP可视化
        self.train_outputs = []
        self.train_labels = []
        self.test_outputs = []
        self.test_labels = []

    def training_step(self, batch, batch_idx):
        x = batch['image']
        y = batch['targets']
        x_hat = self.model(x)
        # 剥离计算图转存CPU,避免占显存
        self.train_outputs.append(x_hat.detach().cpu())
        self.train_labels.append(y.detach().cpu())
        loss = nn.CrossEntropyLoss()(x_hat, y)
        return loss

    def test_step(self, batch, batch_idx):
        x = batch['image']
        y = batch['targets']
        x_hat = self.model(x)
        self.test_outputs.append(x_hat.detach().cpu())
        self.test_labels.append(y.detach().cpu())

    def on_train_epoch_end(self):
        # 拼接所有batch的输出为完整张量
        all_train_outputs = torch.cat(self.train_outputs, dim=0)
        all_train_labels = torch.cat(self.train_labels, dim=0)
        # 此处可添加转存numpy、UMAP可视化的逻辑
        # 清空列表避免下一轮epoch重复存储
        self.train_outputs.clear()
        self.train_labels.clear()

    def on_test_epoch_end(self):
        all_test_outputs = torch.cat(self.test_outputs, dim=0)
        all_test_labels = torch.cat(self.test_labels, dim=0)
        # 此处可添加UMAP可视化逻辑
        self.test_outputs.clear()
        self.test_labels.clear()

    # 以下为Lightning要求必须实现的其他方法,按需补充
    def configure_optimizers(self):
        return torch.optim.Adam(self.parameters(), lr=self.learning_rate)

小提示:如果做特征分布可视化,建议提取模型评估模式下的输出,你可以单独实现predict_step,用trainer.predict()接口跑数据集,模型会自动关闭dropout、冻结BatchNorm参数,得到的特征更稳定。


如何确认输出特征不会被立即删除或丢弃

只要你把x_hat的引用存到了self.xxx这类模型实例的全局属性里,Python的垃圾回收机制就不会回收这块内存,特征不会被丢弃,可通过两种方法验证:

  • 在x_hat = self.model(x)之后打印id(x_hat),epoch结束后打印对应列表里元素的内存地址,两者一致即可证明没有被替换或删除
  • 每跑完一个batch打印len(self.test_outputs),总长度和数据集样本数一致,即可证明所有样本的输出都被正常存储

即使你没有加detach(),反向传播结束后仅会释放x_hat绑定的计算图,x_hat本身的张量值还是会保留,完全不影响可视化使用。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.23 17:36:01