customization 의 sweet spot
처음 손대야 할 칸은 train_step() override 야. 딱 한 method — batch 당 forward/backward 로직 — 만 갈아끼우고 나머지는 전부 상속해. progress bar, callback, validation pass, distribution strategy 가 다 그대로 도는 이유는 loop 를 여전히 fit() 이 몰기 때문이야. 내장 step 대신 네 step 을 호출할 뿐이지.
step 안의 네 동작
모든 train_step() 은 같은 네 가지를 해: batch unpack → forward 로 y_pred → self.compute_loss() 로 loss → gradient 계산 후 apply. 마무리는 metric object update + name → value dict 반환. 이 dict 가 정확히 fit() 이 progress bar 에 찍는 값이야.
backend 가 새는 지점은 gradient mechanics 한 곳뿐 — TF 는 tf.GradientTape, PyTorch 는 loss.backward() + optimizer.step(), JAX 는 stateless compute_loss_and_updates 형태. Keras 3 의 backend 추상화는 layer 단에서 끝나고, training step 의 mechanics 는 backend native 야. 아래는 TF backend 버전.