본문 바로가기
C.W.K.
Stream
Lesson 03 of 06 · published

데이터 파이프라인 선택지

~8 min · data

Level 0Keras 도제
0 XP0/97 lessons0/20 achievements
0/120 XP to next level120 XP to go0% complete

fit()은 여러 입력 형식을 같은 방식으로 받아

Keras 3의 model.fit()은 데이터가 어떤 그릇에 담겼는지와 모델 학습을 분리해. 같은 호출에 NumPy 배열, tf.data.Dataset, PyTorch DataLoader, keras.utils.PyDataset, Pandas DataFrame을 전달할 수 있고 어느 백엔드에서도 사용할 수 있어. 덕분에 모델 코드는 유지한 채 아래쪽 데이터 계층만 교체할 수 있지.

복잡도는 필요한 만큼만

데이터가 모두 RAM에 들어간다면 NumPy가 가장 단순한 선택이야. 메모리를 넘거나 스트리밍, 실행 중 변환, 섞기, 미리 가져오기가 필요할 때 본격적인 파이프라인을 사용해. tf.data는 다중 스레드, 미리 가져오기, 섞기를 폭넓게 지원해. torch.DataLoader는 PyTorch 백엔드에서 자연스럽고 num_workers로 여러 프로세스가 자료를 읽게 할 수 있어. keras.utils.PyDataset는 생성 과정을 직접 제어할 때 쓸 수 있는 프레임워크 중립 선택지이고, JAX에는 고유 데이터 API인 grain이 있어.

배치 크기가 무시되는 흔한 이유

tf.data, DataLoader, PyDataset처럼 이미 배치된 반복 가능 객체를 넘기면 fit()batch_size 인자는 무시돼. 데이터셋이 배치 크기를 이미 결정했기 때문이야. DataLoader(batch_size=64)를 만든 뒤 fit(..., batch_size=32)를 호출해도 실제 배치는 64야. 배치 크기는 실제 배치를 생성하는 계층에서 정해야 해.

Code

여러 입력 형식을 받는 하나의 fit()·python
# All of these work with model.fit():
model.fit(numpy_x, numpy_y)              # NumPy arrays
model.fit(tf_dataset)                     # tf.data.Dataset
model.fit(torch_dataloader)               # PyTorch DataLoader
model.fit(keras_pydataset)                # keras.utils.PyDataset

External links

Exercise

CIFAR-10을 tf.data.Dataset으로 감싸 배치, 섞기, 미리 가져오기를 적용해 학습해. 이어서 KERAS_BACKEND=torch에서 torch.utils.data.DataLoader로 다시 구현하고 에포크 시간을 비교해.

Progress

Progress is local-only — sign in to sync across devices.
이 페이지에서 버그를 발견하셨거나 피드백이 있으세요?문제 신고

댓글 0

🔔 답글 알림 (로그인 필요)
로그인댓글을 남기려면 로그인해 주세요.

아직 댓글이 없어요. 첫 댓글을 남겨보세요.