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

Python에서 TensorFlow Estimator로 대용량 데이터셋 검사하는 방법

타이타닉 데이터셋은 TensorFlow와 Estimator를 활용해 검사할 수 있습니다. 피처(feature)들을 하나씩 순회하면서 리스트 형태로 변환한 뒤, 콘솔에 출력하는 방식입니다.

이 글에서는 순차적 모델을 구축하는 데 유용한 Keras Sequential API를 사용합니다. Sequential 모델은 단순한 레이어 스택 구조로 동작하며, 각 레이어는 정확히 하나의 입력 텐서와 하나의 출력 텐서를 가집니다.

참고로 최소 한 개 이상의 컨볼루션 레이어를 포함하는 신경망은 CNN(Convolutional Neural Network)이라고 부르며, 이 역시 학습 모델을 만드는 데 활용할 수 있습니다.

개발 환경: Google Colaboratory

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

Estimator란 무엇인가?

Estimator는 TensorFlow가 제공하는 완전한 모델의 고수준(high-level) 추상화입니다. 손쉬운 확장성(scaling)과 비동기 학습(asynchronous training)을 지원하도록 설계되었습니다.

여기서는 tf.estimator API를 사용해 로지스틱 회귀(logistic regression) 모델을 학습시킵니다. 이 모델은 다른 알고리즘과 비교하기 위한 베이스라인(baseline)으로 활용됩니다. 타이타닉 데이터셋을 사용하며, 목표는 성별, 나이, 좌석 등급 같은 특성을 바탕으로 승객의 생존 여부를 예측하는 것입니다.

피처 열(Feature Columns)의 역할

Estimator는 피처 열(feature columns)을 통해 원본 입력 피처를 어떻게 해석할지 정의합니다. Estimator는 숫자형 입력 벡터를 기대하므로, 피처 열은 데이터셋의 각 피처를 모델이 이해할 수 있는 형태로 변환하는 방법을 설명해 줍니다. 따라서 적절한 피처 열 조합을 선택하고 활용하는 것이 효과적인 모델을 학습하는 데 필수적입니다.

예제 코드

print("데이터셋을 검사합니다")
ds = make_input_fn(dftrain, y_train, batch_size=10)()
for feature_batch, label_batch in ds.take(1):
    print('일부 피처 키:', list(feature_batch.keys()))
    print()
    print('클래스 배치:', feature_batch['class'].numpy())
    print()
    print('레이블 배치:', label_batch.numpy())

코드 출처: https://www.tensorflow.org/tutorials/estimator/linear

실행 결과

데이터셋을 검사합니다
Some feature keys are: ['sex', 'age', 'n_siblings_spouses', 'parch', 'fare', 'class', 'deck', 'embark_town', 'alone']
A batch of class: [b'First' b'First' b'First' b'Third' b'Third' b'Third' b'First' b'Third'
b'Second' b'Third']
A batch of Labels: [0 1 1 0 0 0 1 0 0 0]

코드 설명

  • 입력 함수(make_input_fn)로 생성된 데이터셋 객체를 검사합니다.
  • ds.take(1)을 통해 배치 하나만 가져와 순회(iterate)합니다.
  • 순회 과정에서 피처 키 목록, 클래스 값 배치, 레이블 배치를 콘솔에 출력합니다.

출력 결과에서 확인할 수 있듯이, 데이터셋에는 sex, age, fare, class, deck 등 다양한 피처가 포함되어 있으며, 각 샘플의 생존 여부는 0 또는 1의 레이블로 표현됩니다.