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

TensorFlow로 꽃 데이터셋 모델 컴파일하고 학습(fit)하는 방법

꽃 데이터셋을 활용한 모델은 compile 메서드로 컴파일하고, fit 메서드로 학습시킬 수 있습니다. fit 메서드에는 훈련 데이터셋(train_ds)과 검증 데이터셋(val_ds)이 매개변수로 전달되며, 학습 반복 횟수인 에포크(epochs) 값도 함께 지정합니다.

데이터셋 소개

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

실행 환경

아래 코드는 구글 코랩(Google Colaboratory)에서 실행됩니다. Google Colab은 브라우저에서 바로 파이썬 코드를 실행할 수 있도록 도와주며, 별도의 환경 설정 없이 GPU(그래픽 처리 장치)를 무료로 사용할 수 있다는 장점이 있습니다. Colaboratory는 Jupyter Notebook을 기반으로 만들어졌습니다.

모델 컴파일 및 학습 코드

print("모델을 컴파일합니다")
model.compile(
    optimizer='adam',
    loss=tf.losses.SparseCategoricalCrossentropy(from_logits=True),
    metrics=['accuracy'])
print("모델을 데이터에 학습시킵니다")
model.fit(
    train_ds,
    validation_data=val_ds,
    epochs=3
)

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

실행 결과

모델을 컴파일합니다
모델을 데이터에 학습시킵니다
Epoch 1/3
92/92 [==============================] - 107s 1s/step - loss: 1.3570 - accuracy: 0.4183 - val_loss: 1.0730 - val_accuracy: 0.5913
Epoch 2/3
92/92 [==============================] - 101s 1s/step - loss: 1.0185 - accuracy: 0.5927 - val_loss: 1.0041 - val_accuracy: 0.6199
Epoch 3/3
92/92 [==============================] - 95s 1s/step - loss: 0.8691 - accuracy: 0.6529 - val_loss: 0.9985 - val_accuracy: 0.6281
<tensorflow.python.keras.callbacks.History at 0x7f2cdcbbba90>

결과 해석

  • 모델의 레이어 구성이 완료되고 데이터 준비가 끝나면, 다음 단계는 compile 메서드를 통해 모델을 컴파일하는 것입니다.
  • 컴파일이 완료되면 fit 메서드를 사용해 입력 데이터셋에 모델을 학습시킵니다.
  • 위 실행 결과를 보면 검증 정확도(val_accuracy)가 훈련 정확도(accuracy)보다 낮게 나타납니다.
  • 이러한 격차는 모델이 과적합(overfitting)되었음을 의미하며, 드롭아웃 적용, 데이터 증강(augmentation), 조기 종료(early stopping) 등의 기법으로 개선할 수 있습니다.