PyTorch Lightning模型获取dense层激活值遇KeyError问题咨询
解决FAtNet模型dense层激活值获取的KeyError问题
问题重现
使用PyTorch Lightning定义的FAtNet模型包含dense线性层,尝试通过forward hook获取该层激活值时触发KeyError: 'dense':
模型中dense层定义:
def __init__(...) ... self.dense = nn.Linear(channels[-1], 64, bias=True) ...
测试代码:
activation = {} def get_activation(name): def hook(model, input, output): activation[name] = output.detach() return hook test_img = cv.imread(f'digimage/100.jpg') test_img = cv.resize(test_img, (128, 128)) test_img = np.moveaxis(test_img, 2, 0) modelftr = load_feature_model(**model_dict) num_ftrs = modelftr.fc.in_features modelftr.fc = torch.nn.Linear(num_ftrs, 228) modelftr.load_state_dict(torch.load('...')) modelftr.dense.register_forward_hook(get_activation('dense')) with torch.no_grad(): modelatt.to('cpu') modelatt.eval() test_img = torch.tensor(test_img).view(-1, 3, 128, 128).float() output = modelcat(test_img) print(activation['dense'])
错误日志:
8 test_img = torch.tensor(test_img).view(-1, 3, 128, 128).float() 9 output = modelcat(test_img) ---> 10 print(activation['dense']) KeyError: 'dense'
错误原因
- Hook注册到了错误的模型实例:代码中将hook注册到
modelftr的dense层,但实际执行推理的是modelcat,导致hook从未被触发。 - 设备不匹配:模型和输入张量可能处于不同设备(如模型在cuda,输入在cpu),或部分模型组件设备不一致,导致hook逻辑未正常执行。
解决方案
1. 确保Hook注册到实际推理的模型
将hook注册到执行推理的modelcat的dense层,而非其他模型实例:
# 替换原注册hook的代码 modelcat.dense.register_forward_hook(get_activation('dense'))
2. 验证模型实例的dense层存在性
先打印模型结构,确认modelcat确实包含dense层:
print(modelcat) # 或遍历子模块确认 for name, module in modelcat.named_modules(): if name == 'dense': print(f"找到dense层: {module}")
3. 统一模型与输入的设备
确保模型和输入张量处于同一设备(cpu或cuda):
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') modelcat.to(device) test_img = torch.tensor(test_img).view(-1, 3, 128, 128).float().to(device)
4. 修正后的完整测试代码
activation = {} def get_activation(name): def hook(model, input, output): activation[name] = output.detach() return hook # 加载并准备测试图像 test_img = cv.imread(f'digimage/100.jpg') test_img = cv.resize(test_img, (128, 128)) test_img = np.moveaxis(test_img, 2, 0) # 初始化并加载目标模型(modelcat) modelcat = FAtNet(image_size=(128,128), in_channels=3, num_blocks=[...], channels=[...]) # 替换为你的模型参数 modelcat.load_state_dict(torch.load('your_model_weights.pth')) modelcat.eval() # 注册hook到modelcat的dense层 modelcat.dense.register_forward_hook(get_activation('dense')) # 统一设备 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') modelcat.to(device) test_img = torch.tensor(test_img).view(-1, 3, 128, 128).float().to(device) # 推理并转换激活值为列表 with torch.no_grad(): output = modelcat(test_img) dense_activation_list = activation['dense'].cpu().numpy().tolist() print(dense_activation_list)
内容的提问来源于stack exchange,提问作者Mohsen Amiri
相关产品推荐
相关产品推荐

