C.W.K.
Stream
Lesson 03 of 07 · published

다중 입력·다중 출력

~8 min · functional

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

입력 하나·출력 하나가 문제 모양이 아닐 때

Functional API 의 강력한 기능 하나 — 서로 다른 데이터 타입을 합치고 여러 예측을 한 model 에서 동시에 내는 거. 실전 문제 중 상당수가 원래부터 다중 입력·다중 출력이야. support ticket 만 봐도 title, body, tags (서로 다른 feature space 세 개) 가 있고, 우선순위(priority) 랑 담당 부서(department) 를 한 번에 예측하고 싶어. 이걸 억지로 단일 입력·단일 출력에 욱여넣으면 model 세 개를 따로 학습하거나, 다 한 벡터로 flatten 해서 구조를 잃어버려.

코드 블록이 표준 모양이야: 입력마다 branch 하나, 각자 layer 로 처리하고 layers.concatenate 로 합친 뒤, 출력마다 head 하나 로 갈라. 입력·출력 전부에 name= 붙이는 게 포인트 — 그래야 데이터랑 loss 를 깨지기 쉬운 positional list 대신 dict 로 넘길 수 있어.

multi-task 는 공짜가 아니야

여러 head 가 trunk 를 공유하니까 각 출력의 gradient 가 같은 merged feature 를 통해 거꾸로 흘러 — 이 공유 표현이 핵심이야 (multi-task model 이 따로 만든 것보다 generalize 잘 되는 이유). 근데 동시에 head 끼리 경쟁해: 한 loss 가 수치적으로 더 크면 학습을 독차지해. 그래서 compileloss_weights 가 있는 거 — head 마다 loss 를 scale 해서 하나가 나머지를 묻어버리지 않게. 동일 weight 로 시작하고, 출력별 loss 지켜보다가 하나가 정체될 때만 재조정해.

Code

입력 3 개, 출력 head 2 개 (support ticket 분류기)·python
# Multi-input, multi-output ticket classification
title_input = keras.Input(shape=(100,), name="title")
body_input = keras.Input(shape=(500,), name="body")
tags_input = keras.Input(shape=(12,), name="tags")

# Process each branch
title_features = layers.Dense(64, activation="relu")(title_input)
body_features = layers.Dense(128, activation="relu")(body_input)
tags_features = layers.Dense(32, activation="relu")(tags_input)

# Merge branches
x = layers.concatenate([title_features, body_features, tags_features])
x = layers.Dense(128, activation="relu")(x)

# Multiple outputs
priority = layers.Dense(3, activation="softmax", name="priority")(x)
department = layers.Dense(5, activation="softmax", name="department")(x)

model = keras.Model(
    inputs=[title_input, body_input, tags_input],
    outputs=[priority, department],
)

External links

Exercise

(image, age) 받아 (class_logits, regression_score) 내는 model 짜. named input/output 사용. 두 loss 로 compile. synthetic data 로 한 epoch 학습.

Progress

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

댓글 0

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

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