register_buffer、nn.Parameter与torch.tensor比较
在一个类中,同样是注册一个 tensor,下面这三个方法有什么不同呢。
self.alpha = torch.tensor(...)
self.register_buffer("alpha", torch.tensor(...))
self.alpha = nn.Parameter(...)
1
2
3
2
3
| 表达方式 | 是否为模型参数 (.parameters()) | 是否会随模型一起保存 (.state_dict()) | 是否会被转移设备(如 .to(device)) |
|---|---|---|---|
self.alpha = torch.tensor(...) | ❌ 否 | ❌ 否 | ❌ 否 |
self.register_buffer("alpha", torch.tensor(...)) | ❌ 否 | ✅ 是 | ✅ 是 |
self.alpha = nn.Parameter(...) | ✅ 是 | ✅ 是 | ✅ 是 |
下面是一个对比测试的小例子:
import torch
import torch.nn as nn
class CompareAlphaWays(nn.Module):
def __init__(self):
super().__init__()
# 普通 tensor,不会被保存或移动
self.alpha_tensor = torch.tensor(0.5)
# buffer,会被保存、移动,但不可训练
self.register_buffer("alpha_buffer", torch.tensor(0.5))
# 参数,会被保存、移动、训练
self.alpha_param = nn.Parameter(torch.tensor(0.5))
def forward(self):
return self.alpha_tensor, self.alpha_buffer, self.alpha_param
model = CompareAlphaWays()
print("\n📦 模型 state_dict:")
for k, v in model.state_dict().items():
print(f"{k}: {v} (device: {v.device})")
print("\n🎯 模型参数列表 (可训练的):")
for name, param in model.named_parameters():
print(f"{name}: {param} (requires_grad: {param.requires_grad})")
print("\n🚀 转移到 GPU(如果可用):")
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model.to(device)
print(f"\n🔍 alpha_tensor.device: {model.alpha_tensor.device} ❌ (不会变)")
print(f"🔍 alpha_buffer.device: {model.alpha_buffer.device} ✅")
print(f"🔍 alpha_param.device: {model.alpha_param.device} ✅")
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
输出说明:
📦 模型 state_dict:
alpha_param: 0.5 (device: cpu)
alpha_buffer: 0.5 (device: cpu)
🎯 模型参数列表 (可训练的):
alpha_param: Parameter containing:
tensor(0.5000, requires_grad=True) (requires_grad: True)
🚀 转移到 GPU(如果可用):
🔍 alpha_tensor.device: cpu ❌ (不会变)
🔍 alpha_buffer.device: cuda:0 ✅
🔍 alpha_param.device: cuda:0 ✅
1
2
3
4
5
6
7
8
9
10
11
12
13
2
3
4
5
6
7
8
9
10
11
12
13
上次更新: 2025/07/21, 10:36:23