Terjemahan disediakan oleh mesin penerjemah. Jika konten terjemahan yang diberikan bertentangan dengan versi bahasa Inggris aslinya, utamakan versi bahasa Inggris.
Fine-tune model pondasi yang tersedia untuk umum dengan ModelTrainer kelas
catatan
Untuk petunjuk tentang menyempurnakan model pondasi di hub pribadi yang dikuratori, lihat. Fine-tune model hub yang dikuratori
Anda dapat menyempurnakan algoritma bawaan atau model pra-terlatih hanya dalam beberapa baris kode menggunakan SDK. SageMaker Python
-
Pertama, temukan ID model untuk model pilihan Anda diModel pondasi yang tersedia.
-
Menggunakan ID model, tentukan pekerjaan pelatihan Anda dengan a 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) -
Panggil
train()metode pada AndaModelTrainer, arahkan ke data pelatihan yang akan digunakan untuk penyetelan halus.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), ] ) -
Kemudian, gunakan
deploymetode untuk secara otomatis menerapkan model Anda untuk inferensi. Dalam contoh ini, kita menggunakan model GPT-J 6B dariHugging Face.from sagemaker.serve import ModelBuilder model_builder = ModelBuilder.from_jumpstart_config(jumpstart_config=jumpstart_config) model = model_builder.build() endpoint = model_builder.deploy() -
Anda kemudian dapat menjalankan inferensi dengan model yang digunakan menggunakan
invokemetode. Text-generation model seperti ini menerima badan permintaan JSON denganinputskunci. Serialkan payload denganjson.dumpsdan atur jenis konten ke.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'))
catatan
Contoh ini menggunakan model dasar GPT-J 6B, yang cocok untuk berbagai kasus penggunaan pembuatan teks termasuk menjawab pertanyaan, pengenalan entitas bernama, ringkasan, dan banyak lagi. Untuk informasi selengkapnya tentang kasus penggunaan model, lihatModel pondasi yang tersedia.
Anda dapat secara opsional menentukan versi model pada AndaJumpStartConfig. Untuk memilih jenis dan jumlah instance, berikan Compute objek keModelTrainer.from_jumpstart_config. Itu JumpStartConfig sendiri tidak menerima pengaturan instance. Untuk informasi selengkapnya tentang ModelTrainer kelas dan parameternya, lihat SageMaker Melatih
Periksa jenis instans default
Saat menyempurnakan model pra-terlatih dengan ModelTrainer kelas, Anda dapat secara opsional menentukan versi model pada kelas Anda. JumpStartConfig Anda juga dapat memilih jenis instance dengan Compute objek. Semua JumpStart model memiliki tipe instance default. Ambil jenis instance pelatihan default menggunakan kode berikut:
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)
Anda dapat melihat semua jenis instance yang didukung untuk JumpStart model tertentu dengan instance_types.retrieve() metode ini.
Periksa hyperparameter default
Untuk memeriksa hyperparameter default yang digunakan untuk pelatihan, Anda dapat menggunakan retrieve_default() metode dari hyperparameters kelas.
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)
Untuk informasi lebih lanjut tentang hyperparameter yang tersedia, lihatHiperparameter fine tuning yang umumnya didukung.
Periksa definisi metrik default
Anda juga dapat memeriksa definisi metrik default:
from sagemaker.core import metric_definitions print(metric_definitions.retrieve_default(model_id=model_id, model_version=model_version))