Функция save
Функция save сохраняет объекты PyTorch (тензоры, модели, словари) в файл на диске. Первым параметром передаётся сохраняемый объект, вторым - путь к файлу, в который будет выполнено сохранение. Функция использует формат сериализации Pickle, но с некоторыми оптимизациями для работы с тензорами.
Синтаксис
torch.save(obj, f)
Пример
Давайте сохраним простой тензор в файл:
import torch
t = torch.tensor([1, 2, 3, 4, 5])
torch.save(t, 'tensor.pt')
После выполнения этого кода в текущей директории появится файл tensor.pt, содержащий сохранённый тензор.
Пример
Давайте сохраним словарь с несколькими тензорами:
import torch
data = {
'tensor1': torch.tensor([1, 2, 3]),
'tensor2': torch.tensor([4, 5, 6]),
'label': 'example'
}
torch.save(data, 'data.pt')
Результат выполнения кода:
"data saved to data.pt"
Пример
Давайте сохраним простую нейросеть:
import torch
import torch.nn as nn
class SimpleModel(nn.Module):
def __init__(self):
super().__init__()
self.fc = nn.Linear(10, 5)
def forward(self, x):
return self.fc(x)
model = SimpleModel()
torch.save(model, 'model.pt')
Результат выполнения кода:
"model saved to model.pt"
Пример
Давайте сохраним только состояние модели (рекомендуемый способ):
import torch
import torch.nn as nn
model = nn.Linear(10, 5)
state_dict = model.state_dict()
torch.save(state_dict, 'model_state.pt')
Результат выполнения кода:
"state dict saved to model_state.pt"
Смотрите также
-
функцию
load,
которая загружает сохранённые объекты из файла -
функцию
jit.save,
которая сохраняет скомпилированную модель в формате TorchScript -
функцию
jit.load,
которая загружает модель из формата TorchScript -
функцию
onnx.export,
которая экспортирует модель в формат ONNX