如何在运行时获取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
相关产品推荐
相关产品推荐

