As traduções são geradas por tradução automática. Em caso de conflito entre o conteúdo da tradução e da versão original em inglês, a versão em inglês prevalecerá.
Fine-tune modelos de fundação disponíveis publicamente com a ModelTrainer classe
nota
Para ver instruções sobre como ajustar os modelos de base em um hub privado selecionado, consulte Fine-tune modelos de hub selecionados.
Você pode ajustar um algoritmo incorporado ou um modelo pré-treinado em apenas algumas linhas de código usando o SDK. SageMaker Python
-
Primeiro, encontre o ID do modelo de sua escolha emModelos de base disponíveis.
-
Usando o ID do modelo, defina seu trabalho de treinamento com um 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) -
Chame o
train()método para vocêModelTrainer, apontando para os dados de treinamento a serem usados para o ajuste fino.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), ] ) -
Em seguida, use o método
deploypara implantar automaticamente seu modelo para inferência. Neste exemplo, usamos o modelo 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() -
Em seguida, você pode executar a inferência com o modelo implantado usando o
invokemétodo. Text-generation modelos como este aceitam um corpo de solicitação JSON com umainputschave. Serialize a carga comjson.dumpse defina o tipo de conteúdo como.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'))
nota
Este exemplo usa o modelo básico GPT-J 6B, que é adequado para uma ampla variedade de casos de uso de geração de texto, incluindo resposta a perguntas, reconhecimento de entidades nomeadas, resumos e muito mais. Para obter mais informações sobre os casos de uso do modelo, consulte Modelos de base disponíveis.
Opcionalmente, você pode especificar uma versão do modelo no seuJumpStartConfig. Para escolher um tipo de instância e contar, passe um Compute objeto paraModelTrainer.from_jumpstart_config. O JumpStartConfig próprio não aceita configurações de instância. Para obter mais informações sobre a ModelTrainer classe e seus parâmetros, consulte SageMaker Treinar
Verificar os tipos de instância padrão
Ao ajustar um modelo pré-treinado com a ModelTrainer classe, você pode, opcionalmente, especificar uma versão do modelo em seu. JumpStartConfig Você também pode escolher um tipo de instância com um Compute objeto. Todos os JumpStart modelos têm um tipo de instância padrão. Recupere o tipo de instância de treinamento padrão usando o seguinte código:
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)
Você pode ver todos os tipos de instância compatíveis com um determinado JumpStart modelo com o instance_types.retrieve() método.
Verificar hiperparâmetros padrão
Para verificar os hiperparâmetros padrão usados para treinamento, você pode usar o método retrieve_default() da função 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)
Para obter mais informações sobre hiperparâmetros disponíveis, consulte Hiperparâmetros de ajuste normalmente aceitos.
Verificar as definições de métricas padrão
Você também pode verificar as definições de métricas padrão:
from sagemaker.core import metric_definitions print(metric_definitions.retrieve_default(model_id=model_id, model_version=model_version))