PyTorch torch.save and torch.load functions
Pytorch torch Reference Manual
torch.saveandtorch.loadare functions in PyTorch for serializing (saving) and deserializing (loading) tensors, models, and other Python objects.
These functions are essential in scenarios such as saving trained models and saving checkpoints to resume training.
Function Definition
torch.save(obj, f, pickle_module, pickle_protocol) torch.load(f, map_location, pickle_module, weights_only)
torch.save Parameters:
obj: The object to be saved, which can be a tensor, model, dictionary, list, etc.f: File path (string or file object).pickle_module(Optional): The module used for serialization.pickle_protocol(Optional): Serialization protocol version.
torch.load Parameters:
f: File path (string or file object).map_location(Optional): Specifies how to map storage to different devices.pickle_module(Optional): The module used for deserialization.weights_only(bool, optional): Whether to load only weights and not Python objects.
Usage Examples
Example 1: Saving and Loading Tensors
Example
import torch
# Create some tensors
x = torch.tensor([1, 2, 3, 4, 5])
y = torch.randn(3, 4)
# Save tensors to file
torch.save({'x': x, 'y': y}, 'tensors.pth')
# Load tensors from file
loaded = torch.load('tensors.pth')
print("Loaded data:", loaded)
print("x:", loaded['x'])
print("y:", loaded['y'])
# Create some tensors
x = torch.tensor([1, 2, 3, 4, 5])
y = torch.randn(3, 4)
# Save tensors to file
torch.save({'x': x, 'y': y}, 'tensors.pth')
# Load tensors from file
loaded = torch.load('tensors.pth')
print("Loaded data:", loaded)
print("x:", loaded['x'])
print("y:", loaded['y'])
The output is:
加载的数据: {'x': tensor([1, 2, 3, 4, 5]), 'y': tensor([[-0.3042, -0.9077, -1.0826, 0.9333],
[ 0.0551, 0.6728, 0.5942, -0.1522],
[-0.3744, 0.9239, -0.2104, -0.5239]])}
x: tensor([1, 2, 3, 4, 5])
y: tensor([[-0.3042, -0.9077, -1.0826, 0.9333],
[ 0.0551, 0.6728, 0.5942, -0.1522],
[-0.3744, 0.9239, -0.2104, -0.5239]])
Example 2: Saving and Loading Models
Example
import torch
import torch.nn as nn
# Define a simple neural network
class SimpleNet(nn.Module):
def __init__(self):
super(SimpleNet, self).__init__()
self.fc1 = nn.Linear(10, 5)
self.fc2 = nn.Linear(5, 2)
def forward(self, x):
x = torch.relu(self.fc1(x))
x = self.fc2(x)
return x
# Create a model instance
model = SimpleNet()
# Save the model (save the entire model)
torch.save(model, 'model.pth')
# Load the model
loaded_model = torch.load('model.pth')
print("Model saved and loaded")
print(loaded_model)
import torch.nn as nn
# Define a simple neural network
class SimpleNet(nn.Module):
def __init__(self):
super(SimpleNet, self).__init__()
self.fc1 = nn.Linear(10, 5)
self.fc2 = nn.Linear(5, 2)
def forward(self, x):
x = torch.relu(self.fc1(x))
x = self.fc2(x)
return x
# Create a model instance
model = SimpleNet()
# Save the model (save the entire model)
torch.save(model, 'model.pth')
# Load the model
loaded_model = torch.load('model.pth')
print("Model saved and loaded")
print(loaded_model)
Example 3: Saving Only Model Parameters (Recommended Method)
Example
import torch
import torch.nn as nn
class SimpleNet(nn.Module):
def __init__(self):
super(SimpleNet, self).__init__()
self.fc1 = nn.Linear(10, 5)
self.fc2 = nn.Linear(5, 2)
def forward(self, x):
x = torch.relu(self.fc1(x))
x = self.fc2(x)
return x
model = SimpleNet()
# Save only model parameters (recommended method)
torch.save(model.state_dict(), 'model_weights.pth')
# Create a new model and load parameters
new_model = SimpleNet()
new_model.load_state_dict(torch.load('model_weights.pth'))
print("Model parameters saved and loaded")
print(new_model.state_dict().keys())
import torch.nn as nn
class SimpleNet(nn.Module):
def __init__(self):
super(SimpleNet, self).__init__()
self.fc1 = nn.Linear(10, 5)
self.fc2 = nn.Linear(5, 2)
def forward(self, x):
x = torch.relu(self.fc1(x))
x = self.fc2(x)
return x
model = SimpleNet()
# Save only model parameters (recommended method)
torch.save(model.state_dict(), 'model_weights.pth')
# Create a new model and load parameters
new_model = SimpleNet()
new_model.load_state_dict(torch.load('model_weights.pth'))
print("Model parameters saved and loaded")
print(new_model.state_dict().keys())
The output is:
模型参数已保存并加载 odict_keys(['fc1.weight', 'fc1.bias', 'fc2.weight', 'fc2.bias'])It is recommended to save only the model parameters (
state_dict), rather than saving the entire model. This allows parameters to be reused between models with different architectures.Example 4: Transferring Between CPU and GPU
Example
import torch
# Assume the model was saved on GPU
if torch.cuda.is_available():
x = torch.randn(2, 3, device='cuda')
torch.save(x, 'tensor_gpu.pth')
# Load on CPU
x_cpu = torch.load('tensor_gpu.pth', map_location='cpu')
print("Loaded to CPU:", x_cpu.device)
Using
map_locationthe parameter, you can load tensors onto different devices.
Notes
- Save files usually use
.pth或.ptextension. - It is recommended to save only the model parameters (
state_dict), rather than saving the entire model. - When loading, security and version compatibility should be considered.