Computer >> 컴퓨터 >  >> 프로그래밍 >> Python

TensorFlow에서 꽃 데이터셋으로 모델 학습을 이어가는 방법

꽃(flower) 데이터셋으로 모델 학습을 계속 진행하려면 Keras의 fit() 메서드를 사용합니다. 이 메서드에 에포크(epoch) 수, 즉 전체 데이터를 몇 번 반복하여 학습할지를 지정할 수 있습니다. 학습이 진행되는 동안 일부 샘플 이미지도 콘솔에 함께 표시됩니다.

꽃 데이터셋 소개

이 튜토리얼에서는 수천 장의 꽃 이미지를 포함하는 꽃 데이터셋을 사용합니다. 데이터셋은 총 5개의 하위 디렉터리로 구성되어 있으며, 각 하위 디렉터리는 하나의 클래스(꽃 종류)에 해당합니다.

아래 코드는 Google Colaboratory(Colab) 환경에서 실행됩니다. Google Colab은 브라우저에서 별도 설정 없이 Python 코드를 실행할 수 있게 해주며, GPU(그래픽 처리 장치)를 무료로 사용할 수 있다는 장점이 있습니다. Colaboratory는 Jupyter Notebook을 기반으로 구축되었습니다.

모델 학습 코드

print("The data is fit to the model")
model.fit(
    train_ds,
    validation_data=val_ds,
    epochs=3
)

코드 출처: https://www.tensorflow.org/tutorials/load_data/images

실행 결과

The data is fit to the model
Epoch 1/3
92/92 [==============================] - 102s 1s/step - loss: 0.7615 - accuracy: 0.7146 - val_loss: 0.7673 - val_accuracy: 0.7180
Epoch 2/3
92/92 [==============================] - 95s 1s/step - loss: 0.5864 - accuracy: 0.7786 - val_loss: 0.6814 - val_accuracy: 0.7629
Epoch 3/3
92/92 [==============================] - 95s 1s/step - loss: 0.4180 - accuracy: 0.8478 - val_loss: 0.7040 - val_accuracy: 0.7575
<tensorflow.python.keras.callbacks.History at 0x7fda872ea940>

코드 설명

  • 기존에 keras.preprocessing으로 만든 데이터셋과 유사한 형태의 데이터셋을 tf.data.Dataset API를 사용해 구축했습니다.
  • 구축된 데이터셋을 그대로 사용하여 모델을 학습시킬 수 있습니다.
  • 학습 시간이 너무 오래 걸리지 않도록 에포크 수를 3으로 제한하여 진행했습니다.

실행 결과를 보면 에포크가 반복될수록 학습 정확도(accuracy)는 0.71 → 0.78 → 0.85로 꾸준히 향상되는 반면, 검증 손실(val_loss)은 세 번째 에포크에서 다소 상승하는 모습을 보입니다. 이는 과적합(overfitting)이 시작되고 있음을 의미할 수 있으므로, 실제 프로젝트에서는 조기 종료(early stopping)나 데이터 증강(augmentation) 같은 기법을 함께 활용하는 것이 좋습니다.