View a markdown version of this page

Cara menggunakan Klasifikasi SageMaker Gambar - TensorFlow algoritma - Amazon SageMaker AI

Terjemahan disediakan oleh mesin penerjemah. Jika konten terjemahan yang diberikan bertentangan dengan versi bahasa Inggris aslinya, utamakan versi bahasa Inggris.

Cara menggunakan Klasifikasi SageMaker Gambar - TensorFlow algoritma

Anda dapat menggunakan Klasifikasi Gambar - TensorFlow sebagai algoritma bawaan Amazon SageMaker AI. Bagian berikut menjelaskan cara menggunakan Klasifikasi Gambar - TensorFlow dengan SageMaker AI Python SDK. Untuk informasi tentang cara menggunakan Klasifikasi Gambar - TensorFlow dari UI Amazon SageMaker Studio Classic, lihatSageMaker JumpStart model pra-terlatih.

Klasifikasi Gambar - TensorFlow algoritma mendukung pembelajaran transfer menggunakan salah satu model TensorFlow Hub pra-latih yang kompatibel. Untuk daftar semua model pra-pelatihan yang tersedia, lihatTensorFlow Model Hub. Setiap model pra-terlatih memiliki keunikan. model_id Contoh berikut menggunakan MobileNet V2 1.00 224 (model_id:tensorflow-ic-imagenet-mobilenet-v2-100-224-classification-4) untuk menyempurnakan dataset kustom. Semua model yang telah dilatih sebelumnya diunduh dari TensorFlow Hub dan disimpan di bucket Amazon S3 sehingga pekerjaan pelatihan dapat berjalan dalam isolasi jaringan. Gunakan artefak pelatihan model yang dibuat sebelumnya ini untuk membangun SageMaker AI ModelTrainer.

Pertama, ambil URI gambar Docker, URI skrip pelatihan, dan URI model yang telah dilatih sebelumnya. Kemudian, ubah hyperparameter sesuai keinginan Anda. Anda dapat melihat kamus Python dari semua hyperparameter yang tersedia dan nilai defaultnya denganhyperparameters.retrieve_default. Untuk informasi selengkapnya, lihat Klasifikasi Gambar - TensorFlow Hyperparameter. Gunakan nilai-nilai ini untuk membangun SageMaker AI ModelTrainer.

catatan

Nilai hyperparameter default berbeda untuk model yang berbeda. Untuk model yang lebih besar, ukuran batch default lebih kecil dan train_only_top_layer hyperparameter diatur ke"True".

Contoh ini menggunakan tf_flowers dataset, yang berisi lima kelas gambar bunga. Kami telah mengunduh kumpulan data dari TensorFlow bawah lisensi Apache 2.0 dan membuatnya tersedia dengan Amazon S3. Untuk menyempurnakan model Anda, hubungi .fit menggunakan lokasi Amazon S3 dari kumpulan data pelatihan Anda.

from sagemaker.core import image_uris from sagemaker.core import model_uris, script_uris, hyperparameters from sagemaker.train import ModelTrainer from sagemaker.train.configs import InputData from sagemaker.train.configs import SourceCode, Compute, StoppingCondition, OutputDataConfig model_id, model_version = "tensorflow-ic-imagenet-mobilenet-v2-100-224-classification-4", "*" training_instance_type = "ml.p3.2xlarge" # Retrieve the Docker image train_image_uri = image_uris.retrieve(model_id=model_id,model_version=model_version,image_scope="training",instance_type=training_instance_type,region=None,framework=None) # Retrieve the training script train_source_uri = script_uris.retrieve(model_id=model_id, model_version=model_version, script_scope="training") # Retrieve the pretrained model tarball for transfer learning train_model_uri = model_uris.retrieve(model_id=model_id, model_version=model_version, model_scope="training") # Retrieve the default hyper-parameters for fine-tuning the model hyperparameters = hyperparameters.retrieve_default(model_id=model_id, model_version=model_version) # [Optional] Override default hyperparameters with custom values hyperparameters["epochs"] = "5" # The sample training data is available in the following S3 bucket training_data_bucket = f"jumpstart-cache-prod-{aws_region}" training_data_prefix = "training-datasets/tf_flowers/" training_dataset_s3_path = f"s3://{training_data_bucket}/{training_data_prefix}" output_bucket = sess.default_bucket() output_prefix = "jumpstart-example-ic-training" s3_output_location = f"s3://{output_bucket}/{output_prefix}/output" # Create SageMaker ModelTrainer instance tf_ic_model_trainer = ModelTrainer( role=aws_role, training_image=train_image_uri, source_code=SourceCode(source_dir=train_source_uri, entry_script="transfer_learning.py"), # In V3, pre-trained model artifacts are passed via input_data_config compute=Compute(instance_type=training_instance_type, instance_count=1), stopping_condition=StoppingCondition(max_runtime_in_seconds=360000), hyperparameters=hyperparameters, output_data_config=OutputDataConfig(s3_output_path=s3_output_location), ) # Use S3 path of the training data to launch SageMaker TrainingJob tf_ic_model_trainer.train( input_data_config=[ InputData(channel_name="training", data_source=training_dataset_s3_path), InputData(channel_name="model", data_source=train_model_uri), ] )