Les traductions sont fournies par des outils de traduction automatique. En cas de conflit entre le contenu d'une traduction et celui de la version originale en anglais, la version anglaise prévaudra.
Fine-tune modèles de base accessibles au public avec la ModelTrainer classe
Note
Pour obtenir des instructions sur le peaufinage des modèles de fondation dans un hub privé organisé, consultez Fine-tune modèles de hubs sélectionnés.
Vous pouvez affiner un algorithme intégré ou un modèle pré-entraîné en quelques lignes de code à l'aide du SageMaker Python SDK.
-
Tout d'abord, recherchez l'identifiant du modèle de votre choix dansModèles de fondation disponibles.
-
À l'aide de l'ID du modèle, définissez votre poste de formation à l'aide d'un JumpStart
ModelTrainer.from sagemaker.train import ModelTrainer from sagemaker.core.jumpstart.configs import JumpStartConfig jumpstart_config = JumpStartConfig(model_id="huggingface-textgeneration1-gpt-j-6b") model_trainer = ModelTrainer.from_jumpstart_config(jumpstart_config=jumpstart_config) -
Appelez la
train()méthode qui vous convientModelTrainer, en pointant vers les données d'entraînement à utiliser pour affiner.from sagemaker.train.configs import InputData model_trainer.train( input_data_config=[ InputData(channel_name="train", data_source=training_dataset_s3_path), InputData(channel_name="validation", data_source=validation_dataset_s3_path), ] ) -
Utilisez ensuite la méthode
deploypour déployer automatiquement votre modèle à des fins d’inférence. Dans cet exemple, nous utilisons le modèle GPT-J 6B deHugging Face.from sagemaker.serve import ModelBuilder model_builder = ModelBuilder.from_jumpstart_config(jumpstart_config=jumpstart_config) model = model_builder.build() endpoint = model_builder.deploy() -
Vous pouvez ensuite exécuter une inférence avec le modèle déployé à l'aide de la
invokeméthode. Text-generation les modèles comme celui-ci acceptent un corps de requête JSON avec uneinputsclé. Sérialisez la charge utile avecjson.dumpset définissez le type de contenu sur.application/jsonimport json question ="What is Southern California often abbreviated as?"payload = {"inputs": question, "parameters": {"max_new_tokens": 100}} response = endpoint.invoke(body=json.dumps(payload), content_type="application/json") print(response.body.read().decode('utf-8'))
Note
Cet exemple utilise le modèle de base GPT-J 6B, qui convient à un large éventail de cas d'utilisation de génération de texte, notamment la réponse à des questions, la reconnaissance d'entités nommées, la synthèse, etc. Pour plus d’informations sur les cas d’utilisation d’un modèle, consultez Modèles de fondation disponibles.
Vous pouvez éventuellement spécifier une version de modèle sur votreJumpStartConfig. Pour choisir un type et un nombre d'instances, transmettez un Compute objet àModelTrainer.from_jumpstart_config. Le JumpStartConfig lui-même n'accepte pas les paramètres d'instance. Pour plus d'informations sur la ModelTrainer classe et ses paramètres, consultez la documentation SageMaker Train
Vérification de types d’instance par défaut
Lorsque vous peaufinez un modèle pré-entraîné avec la ModelTrainer classe, vous pouvez éventuellement spécifier une version du modèle sur votre. JumpStartConfig Vous pouvez également choisir un type d'instance avec un Compute objet. Tous les JumpStart modèles possèdent un type d'instance par défaut. Extrayez le type d’instance d’entraînement par défaut à l’aide du code suivant :
from sagemaker.core import instance_types instance_type = instance_types.retrieve_default( model_id=model_id, model_version=model_version, scope="training") print(instance_type)
Vous pouvez voir tous les types d'instances pris en charge pour un JumpStart modèle donné à l'aide de la instance_types.retrieve() méthode.
Vérification d’hyperparamètres par défaut
Pour vérifier les hyperparamètres par défaut utilisés pour l’entraînement, vous pouvez utiliser la méthode retrieve_default() de la classe hyperparameters.
from sagemaker.core import hyperparameters my_hyperparameters = hyperparameters.retrieve_default(model_id=model_id, model_version=model_version) print(my_hyperparameters) # Optionally override default hyperparameters for fine-tuning my_hyperparameters["epoch"] = "3" my_hyperparameters["per_device_train_batch_size"] = "4" # Optionally validate hyperparameters for the model hyperparameters.validate(model_id=model_id, model_version=model_version, hyperparameters=my_hyperparameters)
Pour plus d’informations sur les hyperparamètres disponibles, consultez Hyperparamètres de peaufinage couramment pris en charge.
Vérification de définitions de métriques par défaut
Vous pouvez également vérifier les définitions de métriques par défaut :
from sagemaker.core import metric_definitions print(metric_definitions.retrieve_default(model_id=model_id, model_version=model_version))