PyTorch中nn.Module示例詳解
直接print(dir(nn.Module)),得到如下內容:

一、模型結構與參數
parameters()- 用途:返回模塊的所有可訓練參數(如權重、偏置)。
- 示例:
for param in model.parameters(): print(param.shape)
named_parameters()- 用途:返回帶名稱的參數迭代器,便于調試和訪問特定參數。
- 示例:
for name, param in model.named_parameters(): if 'weight' in name: print(name, param.shape)
children()- 用途:返回直接子模塊的迭代器。
- 示例:
for child in model.children(): print(type(child))
modules()- 用途:遞歸返回所有子模塊(包括自身)。
- 示例:
for module in model.modules(): if isinstance(module, nn.Conv2d): print(module.kernel_size)
二、模型狀態(tài)與模式
train()和eval()- 用途:切換訓練/推理模式(影響Dropout、BatchNorm等層)。
- 示例:
model.train() # 訓練模式 model.eval() # 推理模式
training- 用途:布爾屬性,指示當前模式(
True為訓練,False為推理)。 - 示例:
print(model.training) # 輸出:True/False
- 用途:布爾屬性,指示當前模式(
三、模型保存與加載
state_dict()- 用途:返回包含模型所有參數的字典(
OrderedDict)。 - 示例:
torch.save(model.state_dict(), 'model.pth')
- 用途:返回包含模型所有參數的字典(
load_state_dict()- 用途:從字典加載模型參數。
- 示例:
model.load_state_dict(torch.load('model.pth'))
四、設備與數據類型
to()- 用途:將模型移動到指定設備(如GPU)或轉換數據類型。
- 示例:
model.to('cuda') # 移動到GPU model.to(torch.float16) # 轉換為半精度
cpu()和cuda()- 用途:快捷方法,分別將模型移動到CPU或GPU。
- 示例:
model.cuda() # 等價于 model.to('cuda')
五、前向傳播與計算
forward()- 用途:定義模型的前向傳播邏輯(需在自定義模塊中重寫)。
- 示例:
class MyModel(nn.Module): def forward(self, x): return self.layer(x)
__call__()- 用途:調用模型實例時觸發(fā)(內部調用
forward(),支持鉤子函數)。 - 示例:
output = model(x) # 等價于 output = model.forward(x)
- 用途:調用模型實例時觸發(fā)(內部調用
六、參數初始化與優(yōu)化
zero_grad()- 用途:清空所有參數的梯度(通常在每個訓練步驟前調用)。
- 示例:
optimizer.zero_grad() # 等價于 model.zero_grad()
requires_grad_()- 用途:設置參數是否需要梯度(用于凍結部分模型)。
- 示例:
for param in model.parameters(): param.requires_grad = False # 凍結所有參數
七、調試與信息
extra_repr()- 用途:自定義模塊打印信息(需在子類中重寫)。
- 示例:
class MyModel(nn.Module): def extra_repr(self): return f"hidden_size={self.hidden_size}"
dump_patches()- 用途:打印模型的補丁信息(用于調試版本差異)。
八、其他實用方法
apply()- 用途:遞歸應用函數到所有子模塊(如初始化權重)。
- 示例:
def init_weights(m): if isinstance(m, nn.Conv2d): nn.init.kaiming_normal_(m.weight) model.apply(init_weights)
register_forward_hook()- 用途:注冊前向傳播鉤子(用于捕獲中間輸出,調試或特征提?。?。
總結
日常使用中,最頻繁的方法包括:
- 模型構建:
parameters(),children(),modules() - 訓練與推理:
train(),eval(),zero_grad(),forward() - 保存與加載:
state_dict(),load_state_dict() - 設備管理:
to(),cuda(),cpu()
其他方法根據具體需求選擇使用,例如鉤子函數用于高級調試,apply() 用于統(tǒng)一初始化。
與nn.Sequential對比:
1. 繼承關系與基礎屬性
nn.Module- 是所有神經網絡模塊的基類,提供最基礎的功能(如參數管理、鉤子機制)。
- 包含核心屬性:
_parameters,_modules,_buffers等。
nn.Sequential- 是
nn.Module的子類,繼承了所有基礎功能。 - 額外添加了與順序執(zhí)行相關的屬性(如
__getitem__、append)。
- 是
2. 核心差異對比
| 功能類別 | nn.Module | nn.Sequential |
|---|---|---|
| 模塊構建 | 需要手動實現 forward 方法 | 自動按順序執(zhí)行子模塊,無需定義 forward |
| 子模塊訪問 | 通過屬性名(如 self.conv1) | 通過索引或命名(如 model[0]) |
| 動態(tài)修改 | 需手動管理子模塊 | 支持 append、extend、insert 等操作 |
| 適用場景 | 復雜網絡結構(如ResNet、U-Net) | 簡單順序結構(如LeNet卷積部分) |
3. 具體方法對比
3.1 公共方法(兩者都有)
# 模型參數與結構 ['parameters', 'named_parameters', 'children', 'modules', 'named_children', 'named_modules'] # 模型狀態(tài) ['train', 'eval', 'training', 'zero_grad', 'requires_grad_'] # 設備與數據類型 ['to', 'cpu', 'cuda', 'float', 'double', 'half', 'bfloat16'] # 保存與加載 ['state_dict', 'load_state_dict'] # 鉤子機制 ['register_forward_hook', 'register_backward_hook']
3.2nn.Sequential特有的方法
# 列表操作(動態(tài)修改模塊順序) ['__getitem__', '__setitem__', '__delitem__', '__len__', 'append', 'extend', 'insert', 'pop'] # 索引相關 ['_get_item_by_idx']
3.3nn.Module特有的方法
# 自定義實現 ['forward', 'extra_repr'] # 高級管理 ['add_module', 'register_module', 'register_parameter', 'register_buffer']
4. 示例對比
4.1 創(chuàng)建模型
# nn.Module(需自定義 forward)
class CustomModel(nn.Module):
def __init__(self):
super().__init__()
self.conv = nn.Conv2d(3, 64, 3)
self.relu = nn.ReLU()
def forward(self, x):
return self.relu(self.conv(x))
# nn.Sequential(自動按順序執(zhí)行)
seq_model = nn.Sequential(
nn.Conv2d(3, 64, 3),
nn.ReLU()
)4.2 訪問子模塊
# nn.Module custom_model.conv # 通過屬性名訪問 # nn.Sequential seq_model[0] # 通過索引訪問 seq_model.append(nn.MaxPool2d(2)) # 動態(tài)添加模塊
5. 總結
| 特性 | nn.Module | nn.Sequential |
|---|---|---|
| 靈活性 | 高(自定義任意邏輯) | 低(僅支持順序執(zhí)行) |
| 代碼復雜度 | 較高(需手動實現 forward) | 低(自動處理前向傳播) |
| 動態(tài)修改 | 不支持直接操作(需手動管理) | 支持 append、insert 等操作 |
| 適用場景 | 復雜網絡、分支結構、自定義操作 | 簡單堆疊模塊(如CNN的卷積部分) |
建議:
- 對于簡單的順序網絡,優(yōu)先使用
nn.Sequential以減少代碼量。 - 對于包含復雜邏輯(如殘差連接、多輸入輸出)的網絡,使用
nn.Module自定義實現。
到此這篇關于PyTorch中nn.Module詳解的文章就介紹到這了,更多相關PyTorch nn.Module內容請搜索腳本之家以前的文章或繼續(xù)瀏覽下面的相關文章希望大家以后多多支持腳本之家!
相關文章
基于np.arange與np.linspace細微區(qū)別(數據溢出問題)
這篇文章主要介紹了基于np.arange與np.linspace細微區(qū)別(數據溢出問題),具有很好的參考價值,希望對大家有所幫助。如有錯誤或未考慮完全的地方,望不吝賜教2022-05-05

