PyTorch torch.load_state_dict Function


Pytorch torch 参考手册Pytorch torch Reference Manual

torch.load_state_dictIt is a function in PyTorch used to load state dictionaries. It loads the state of model optimizers, etc., from a dictionary object.

Function Definition

torch.load_state_dict(state_dict, strict=True)

Usage Example

Example

import torch
import torch.nn as nn

# Define a simple model
model = nn.Linear(10, 5)

# Save the model state dictionary
state_dict = model.state_dict()
torch.save(state_dict, 'model_state.pt')

# Load the state dictionary
loaded_state = torch.load('model_state.pt')
model.load_state_dict(loaded_state)
print("State dict loaded successfully!")

Pytorch torch 参考手册Pytorch torch Reference Manual

Other Extensions