PyTorch torch.save and torch.load functions


Pytorch torch 参考手册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'])

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)

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())

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)

Usingmap_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.

Pytorch torch 参考手册Pytorch torch Reference Manual