Sakana AI, 역전파 없이 1000층 신경망을 학습하는 PC-ALM 공개
- Sakana AI의 PC-ALM은 역전파 없이 최대 1000층 residual MLP를 학습하며, 층별 지역 동역학만 써서 backprop 성능에 거의 근접함
- 각 층에 dual 뉴런(Lagrange 승수)을 붙여 층의 국소 재귀를 PI 피드백 제어기로 만들었고, 선형 네트워크 극한에서 이 뉴런들이 정확한 backprop credit 신호로 수렴함
- 표준 predictive coding의 신호 감쇠 문제를 넘어섰고, PC가 특히 취약한 깊고 좁은 네트워크 영역을 실험 대상으로 삼음
- 실험은 Fashion-MNIST, CIFAR-10 같은 단순 과제와 residual MLP에 한정되며, 댓글에 따르면 MNIST 정확도는 약 85% 수준임
- 논문은 arXiv 2605.31022, 코드는 github.com/SakanaAI/pc-alm로 공개. 동기는 뇌의 multilayer credit assignment 이해와 뉴로모픽 하드웨어의 에너지 효율임
Hacker News opinions
MNIST 85%라니. CIFAR-10은 그럼 얼마나 나오는데? ImageNet은? 흥미로운 연구인 건 맞는데 이걸 backprop 대안이라고 부르긴 좀 그렇다
MNIST만 있는 게 아니라 이미지 분류 벤치마크에 둘 다 들어있음. 논문 사이트 figure 링크 보면 나옴
얘네가 backprop을 대체하는 게 목표라고 안 했음. backprop을 못 하는 분산 시스템, 그러니까 뇌가 어떻게 학습하는지 이해하는 게 목적이라고 직접 써놨더라
뇌과학 쪽 함의가 진짜 흥미롭다. Friston의 Markov blanket, free energy principle이랑 연결되는 모델 아닌가 싶음
FEP도 결국 KL divergence에 대한 loss 최소화 아님? 우리가 이미 수년간 ML에서 해온 거랑 같은 걸 이름만 바꾼 거 같은데, 차이가 뭔지 모르겠음
credit assignment를 푸는 대안으로 predictive coding 연구가 많은데, 임의의 computation graph에서 backprop과 똑같은 gradient가 나온다는 논문이 특히 좋았음. MIT NEUCO에 실린 'Predictive Coding Approximates Backprop Along Arbitrary Computation Graphs'
그 논문은 fixed prediction assumption을 쓴 거라 좀 혼란을 만들었음. PC가 아니라 반창고 붙인 PC임. 그래도 Beren Millidge 논문들은 다 괜찮고, 최근 Dwarkesh 팟캐스트에도 나왔더라
뇌가 쓰는 학습 알고리즘이 마법은 아니라는 강한 신호였음. backprop보다 compute나 data 효율은 떨어져도, 국소적이고 분산된 방식으로 구현하기는 훨씬 쉽다는 얘기임. 뇌한테는 활성값 저장이나 전역 연결이 엄청 비싸니까
이거 continual learning이랑 관련 있나? 전체 시스템을 멈추지 않고 업데이트할 수 있다는 거잖아
backprop으로도 추론과 업데이트를 동시에 못 할 이유는 없음. 학습 스텝이 pretraining이냐 finetuning이냐 같은 건 수학적으로 같은 연산이고 의도만 다름
catastrophic forgetting을 막는 성질이 있을 수도 있고, 일단은 생물학적으로 그럴듯한 backprop 근사라서 학습 compute를 줄이는 데 쓸모가 있을 듯
continual learning은 weight 단위 비즈니스 모델 문제가 걸림. 랩들 이미 적자라 shared weight를 포기할 리 없고, I/O랑 스토리지 비용 때문에 불가능함. KV cache도 일종의 동적 fast weight인데 그것도 비쌈
backprop으로 학습된 LLM을 이 방식으로 finetuning하면 메모리랑 compute가 덜 들까? LoRA와 full finetuning 사이에 새로운 스펙트럼이 생기는 셈인데
Lagrangian이 시간 평활된 최적화 방향 상태처럼 동작하는데, 미분 불가능한 경계에도 심을 수 있다는 게 핵심임. argmax, MoE, VQ-VAE 같은 이산 컴포넌트 학습에도 쓸 수 있지 않을까
궁금한 게 pipeline parallelism의 타이밍 순서 요구를 완화하는 데 쓸 수 있나? 노드가 통신하면서 lambda를 업데이트하고 다른 시점에 weight를 최적화하는 식으로. batch랑 T를 같은 타임라인에 놓고 descent 방향을 유지할 수도 있고