Estaba buscando formas alternativas de guardar un modelo entrenado en PyTorch. Hasta ahora, he encontrado dos alternativas.
Me he encontrado con esta discusión donde se recomienda el enfoque 2 sobre el enfoque 1.
Mi pregunta es, ¿por qué se prefiere el segundo enfoque? ¿Es solo porque los módulos de torch.nn tienen esas dos funciones y se nos anima a usarlos?
Encontré esta página en su repositorio de github, simplemente copiaré y pegaré el contenido aquí.
Hay dos enfoques principales para serializar y restaurar un modelo.
El primero (recomendado) guarda y carga solo los parámetros del modelo:
torch.save(the_model.state_dict(), PATH)Entonces despúes:
the_model = TheModelClass(*args, **kwargs) the_model.load_state_dict(torch.load(PATH))El segundo guarda y carga todo el modelo:
torch.save(the_model, PATH)Entonces despúes:
the_model = torch.load(PATH)Sin embargo, en este caso, los datos serializados están vinculados a las clases específicas y la estructura de directorio exacta utilizada, por lo que pueden romperse de varias maneras cuando se usan en otros proyectos o después de algunas refactorizaciones serias.
Actualización : consulte también la sección Guardar y cargar el modelo del tutorial de PyTorch
La biblioteca pickle Python implementa protocolos binarios para serializar y deserializar un objeto de Python.
Cuando import torch (o cuando usa PyTorch), import pickle por usted y no necesita llamar a pickle.dump() y pickle.load() directamente, que son los métodos para guardar y cargar el objeto.
De hecho, torch.save() y torch.load() envolverán pickle.dump() y pickle.load() por ti.
Un state_dict la otra respuesta mencionada merece solo algunas notas más.
¿Qué state_dict tenemos dentro de PyTorch? En realidad, hay dos state_dict s.
El modelo de PyTorch es torch.nn.Module que tiene una llamada model.parameters() para obtener parámetros de aprendizaje (w y b). Estos parámetros de aprendizaje, una vez establecidos aleatoriamente, se actualizarán con el tiempo a medida que aprendamos. Los parámetros que se pueden aprender son el primer state_dict .
El segundo state_dict es el dictado de estado del optimizador. Recuerda que el optimizador se utiliza para mejorar nuestros parámetros de aprendizaje. Pero el optimizador state_dict está arreglado. Nada que aprender allí.
Debido a que los objetos state_dict son diccionarios de Python, se pueden guardar, actualizar, modificar y restaurar fácilmente, lo que agrega una gran cantidad de modularidad a los modelos y optimizadores de PyTorch.
Vamos a crear un modelo súper simple para explicar esto:
import torch import torch.optim as optim model = torch.nn.Linear(5, 2) # Initialize optimizer optimizer = optim.SGD(model.parameters(), lr=0.001, momentum=0.9) print("Model's state_dict:") for param_tensor in model.state_dict(): print(param_tensor, "\t", model.state_dict()[param_tensor].size()) print("Model weight:") print(model.weight) print("Model bias:") print(model.bias) print("---") print("Optimizer's state_dict:") for var_name in optimizer.state_dict(): print(var_name, "\t", optimizer.state_dict()[var_name])Este código generará lo siguiente:
Model's state_dict: weight torch.Size([2, 5]) bias torch.Size([2]) Model weight: Parameter containing: tensor([[ 0.1328, 0.1360, 0.1553, -0.1838, -0.0316], [ 0.0479, 0.1760, 0.1712, 0.2244, 0.1408]], requires_grad=True) Model bias: Parameter containing: tensor([ 0.4112, -0.0733], requires_grad=True) --- Optimizer's state_dict: state {} param_groups [{'lr': 0.001, 'momentum': 0.9, 'dampening': 0, 'weight_decay': 0, 'nesterov': False, 'params': [140695321443856, 140695321443928]}]Tenga en cuenta que este es un modelo mínimo. Puede intentar agregar una pila de secuencias
model = torch.nn.Sequential( torch.nn.Linear(D_in, H), torch.nn.Conv2d(A, B, C) torch.nn.Linear(H, D_out), ) Tenga en cuenta que solo las capas con parámetros que se pueden aprender (capas convolucionales, capas lineales, etc.) y los búferes registrados (capas de normas por lotes) tienen entradas en el state_dict del modelo.
Las cosas que no se pueden aprender pertenecen al objeto del optimizador state_dict , que contiene información sobre el estado del optimizador, así como los hiperparámetros utilizados.
El resto de la historia es la misma; en la fase de inferencia (esta es una fase en la que usamos el modelo después del entrenamiento) para predecir; predecimos basándonos en los parámetros que aprendimos. Entonces, para la inferencia, solo necesitamos guardar los parámetros model.state_dict() .
torch.save(model.state_dict(), filepath)Y para usar más tarde model.load_state_dict(torch.load(filepath)) model.eval()
Nota: No olvide la última línea model.eval() esto es crucial después de cargar el modelo.
Tampoco intente guardar torch.save(model.parameters(), filepath) . El model.parameters() es solo el objeto generador.
Por otro lado, torch.save(model, filepath) guarda el objeto del modelo en sí, pero tenga en cuenta que el modelo no tiene el state_dict del optimizador. Verifique la otra excelente respuesta de @Jadiel de Armas para guardar el dictado de estado del optimizador.
Depende de lo que quieras hacer.
Caso # 1: Guarde el modelo para usarlo usted mismo para la inferencia : Guarda el modelo, lo restaura y luego cambia el modelo al modo de evaluación. Esto se hace porque generalmente tiene capas BatchNorm y Dropout que, de manera predeterminada, están en modo de entrenamiento en construcción:
torch.save(model.state_dict(), filepath) #Later to restore: model.load_state_dict(torch.load(filepath)) model.eval()Caso n.º 2: guarde el modelo para reanudar el entrenamiento más tarde : si necesita seguir entrenando el modelo que está a punto de guardar, debe guardar más que solo el modelo. También debe guardar el estado del optimizador, las épocas, la puntuación, etc. Lo haría así:
state = { 'epoch': epoch, 'state_dict': model.state_dict(), 'optimizer': optimizer.state_dict(), ... } torch.save(state, filepath) Para reanudar el entrenamiento, haría cosas como: state = torch.load(filepath) , y luego, para restaurar el estado de cada objeto individual, algo como esto:
model.load_state_dict(state['state_dict']) optimizer.load_state_dict(state['optimizer']) Dado que está reanudando el entrenamiento, NO llame a model.eval() una vez que restaure los estados al cargar.
Caso #3: Modelo para ser usado por otra persona sin acceso a tu código : En Tensorflow puedes crear un archivo .pb que define tanto la arquitectura como los pesos del modelo. Esto es muy útil, especialmente cuando se usa Tensorflow serve . La forma equivalente de hacer esto en Pytorch sería:
torch.save(model, filepath) # Then later: model = torch.load(filepath)Esta forma todavía no es a prueba de balas y dado que pytorch todavía está experimentando muchos cambios, no lo recomendaría.
Una convención común de PyTorch es guardar modelos usando una extensión de archivo .pt o .pth.
Guardar/Cargar todo el modelo
Salvar:
path = "username/directory/lstmmodelgpu.pth" torch.save(trainer, path)Carga:
(La clase de modelo debe definirse en alguna parte)
model.load_state_dict(torch.load(PATH)) model.eval()Si desea guardar el modelo y desea reanudar el entrenamiento más tarde:
GPU única: Guardar:
state = { 'epoch': epoch, 'state_dict': model.state_dict(), 'optimizer': optimizer.state_dict(), } savepath='checkpoint.t7' torch.save(state,savepath)Carga:
checkpoint = torch.load('checkpoint.t7') model.load_state_dict(checkpoint['state_dict']) optimizer.load_state_dict(checkpoint['optimizer']) epoch = checkpoint['epoch']GPU múltiple: Guardar
state = { 'epoch': epoch, 'state_dict': model.module.state_dict(), 'optimizer': optimizer.state_dict(), } savepath='checkpoint.t7' torch.save(state,savepath)Carga:
checkpoint = torch.load('checkpoint.t7') model.load_state_dict(checkpoint['state_dict']) optimizer.load_state_dict(checkpoint['optimizer']) epoch = checkpoint['epoch'] #Don't call DataParallel before loading the model otherwise you will get an error model = nn.DataParallel(model) #ignore the line if you want to load on Single GPU La forma en que guarde su modelo depende de cómo desee acceder a él en el futuro. Si puede llamar a una nueva instancia de la clase del model , entonces todo lo que necesita hacer es guardar/cargar los pesos del modelo con model.state_dict() :
# Save: torch.save(old_model.state_dict(), PATH) # Load: new_model = TheModelClass(*args, **kwargs) new_model.load_state_dict(torch.load(PATH)) Si no puede por alguna razón (o prefiere la sintaxis más simple), entonces puede guardar el modelo completo (en realidad, una referencia a los archivos que definen el modelo, junto con su state_dict) con torch.save() :
# Save: torch.save(old_model, PATH) # Load: new_model = torch.load(PATH)Pero dado que esta es una referencia a la ubicación de los archivos que definen la clase del modelo, este código no es portátil a menos que esos archivos también se transfieran a la misma estructura de directorios.
Si desea que su modelo sea portátil, puede importarlo fácilmente con torch.hub . Si agrega un archivo hubconf.py adecuadamente definido a un repositorio de github, se puede llamar fácilmente desde PyTorch para permitir que los usuarios carguen su modelo con/sin pesos:
hubconf.py ( github.com/repo_propietario/repo_nombre )
dependencies = ['torch'] from my_module import mymodel as _mymodel def mymodel(pretrained=False, **kwargs): return _mymodel(pretrained=pretrained, **kwargs)Cargando modelo:
new_model = torch.hub.load('repo_owner/repo_name', 'mymodel') new_model_pretrained = torch.hub.load('repo_owner/repo_name', 'mymodel', pretrained=True)pip install pytorch-relámpago
asegúrese de que su modelo principal use pl.LightningModule en lugar de nn.Module
Guardando y cargando puntos de control usando pytorch lightning
import pytorch_lightning as pl model = MyLightningModule(hparams) trainer.fit(model) trainer.save_checkpoint("example.ckpt") new_model = MyModel.load_from_checkpoint(checkpoint_path="example.ckpt")En estos días todo está escrito en el tutorial oficial: https://pytorch.org/tutorials/beginner/saving_loading_models.html
Tiene varias opciones sobre cómo guardar y qué guardar y todo se explica en ese tutorial.
Siempre prefiero usar Torch7 (.t7) o Pickle (.pth, .pt) para guardar los pesos de los modelos de pytorch.