Empresas
Empleos
  • Sobre nosotros
  • Soluciones
    • Publicación de vacantes
      Publica tu vacante y recibe candidatos calificados en 48h.
    • Evaluación de candidatos
      500+ pruebas técnicas y psicológicas, más anti-fraude.
    • Headhunting
      Búsqueda ejecutiva a la medida de principio a fin.
    • Nómina + EOR
      Dispersión de nómina y EOR en más de 15 países de LATAM.
  • Precios
  • Empleos

0

1.2K
Vistas
¿La mejor manera de guardar un modelo entrenado en PyTorch?

Estaba buscando formas alternativas de guardar un modelo entrenado en PyTorch. Hasta ahora, he encontrado dos alternativas.

  1. torch.save() para guardar un modelo y torch.load() para cargar un modelo.
  2. model.state_dict() para guardar un modelo entrenado y model.load_state_dict() para cargar el modelo guardado.

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?

over 4 years ago · Santiago Trujillo
9 Respuestas
Responde la pregunta

0

Encontré esta página en su repositorio de github, simplemente copiaré y pegaré el contenido aquí.


Enfoque recomendado para guardar un modelo

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

over 4 years ago · Santiago Trujillo Denunciar

0

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.

over 4 years ago · Santiago Trujillo Denunciar

0

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.

over 4 years ago · Santiago Trujillo Denunciar

0

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()
over 4 years ago · Santiago Trujillo Denunciar

0

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
over 4 years ago · Santiago Trujillo Denunciar

0

Guardar localmente

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.

Guardar en la nube - TorchHub

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)
over 4 years ago · Santiago Trujillo Denunciar

0

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")
over 4 years ago · Santiago Trujillo Denunciar

0

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.

over 4 years ago · Santiago Trujillo Denunciar

0

Siempre prefiero usar Torch7 (.t7) o Pickle (.pth, .pt) para guardar los pesos de los modelos de pytorch.

over 4 years ago · Santiago Trujillo Denunciar
Responde la pregunta
Encuentra empleos remotos

¡Descubre la nueva forma de encontrar empleo!

Top de empleos
Top categorías de empleo
Empresas
Publicar vacante Precios Comercial
Legal
Términos y condiciones Política de privacidad
© 2026 PeakU Inc. All Rights Reserved.
Andres GPT
Recomiéndame algunas ofertas
Necesito ayuda