6.1 KiB
JAX PPO Model-Capacity Diagnostic - 2026-07-12
"MLP를 더 키우면 아직 얻을 게 남아 있나?"에 대한 진단이다.
신규 학습 없이 기존 체크포인트, league_metrics.jsonl, 소스만 읽어 작성했다.
새로 실행한 학습·평가·벤치마크는 없다.
Subject: final candidate
/mnt/2tbhdd/coolrl-lost-cities-artifacts/final-cycles/2026-07-05/final_candidate
Metrics: .../final-cycles/2026-07-05/league/2026-07-05_191933_jax-ppo-final-cycle-c01/league_metrics.jsonl (500 updates)
요약
용량이 현재의 병목일 가능성은 낮다. 다만 PPO 스택에서 모델 크기는 한 번도 검증된 적이 없으므로, 이 판단은 아직 측정으로 뒷받침되지 않았다. 단발 A/B로 결론을 확정할 것을 권한다 — docs/plans/jax-ppo-model-size-ab.md.
1. 현재 크기는 근거 없는 상속값
configs/jax_ppo/ 의 실 학습 config는 전부 hidden_size: 512, num_layers: 3
이다 (smoke.yaml만 64×2). final candidate도 동일하다
(main_ppo_config.json → network: {hidden_size: 512, num_layers: 3}).
hidden_size를 다룬 기존 문서는 전부 은퇴한 Deep CFR/PyTorch 스택 것이다
(input_dim=365, configs/deep_cfr/ 기준). 특히
docs/plans/model_size_experiment.md는
현 JAX PPO 스택과 무관하다.
→ 512×3은 실험으로 고른 값이 아니라 구 스택에서 복사돼 온 값이다.
2. 데이터 대비 모델이 매우 작다 (스케일링에 유리한 신호)
| 항목 | 값 |
|---|---|
OBS_DIM (types.py:27) |
454 |
N_ACTIONS |
96 |
| 512×3 파라미터 수 | 808,033 |
| 1024×4 파라미터 수 | 3,714,145 |
| 2048×4 파라미터 수 | 13,719,649 |
| 최종 사이클 소비 transition | 1,638,400,000 |
| 샘플 : 파라미터 | 약 2,028 : 1 |
transition 수는 rollout_steps=400 × batch_games=8192 × updates=500.
셀프플레이라 데이터는 사실상 무한하다. 샘플:파라미터 2,000:1은 데이터 제약이 아니라 용량/최적화 제약 구간의 전형이며, 보통 이런 영역에서 스케일링이 먹힌다. 0.8M 파라미터는 절대적으로도 작다.
여기까지만 보면 "키우면 이득"이다. 그러나 3절이 반대 방향을 가리킨다.
3. 학습 로그에 용량 부족의 지문이 없다
최종 사이클 500 업데이트 구간 평균:
| update | value_loss | entropy_mean | approx_kl | return_mean | return_std |
|---|---|---|---|---|---|
| 0–50 | 0.0468 | 1.1159 | 0.0433 | 0.2696 | 0.6299 |
| 50–100 | 0.0462 | 1.1167 | 0.0438 | 0.2799 | 0.6275 |
| 100–200 | 0.0461 | 1.1241 | 0.0441 | 0.2822 | 0.6316 |
| 200–300 | 0.0459 | 1.1143 | 0.0437 | 0.2857 | 0.6306 |
| 300–400 | 0.0460 | 1.1239 | 0.0424 | 0.2900 | 0.6323 |
| 400–500 | 0.0463 | 1.1163 | 0.0423 | 0.2916 | 0.6340 |
세 가지를 읽을 수 있다.
(a) 크리틱은 이미 잘 맞춘다. value_loss는 순수 MSE다
(ppo.py:624: masked_mean((mb_returns - value) ** 2, active_weight)).
return_std ≈ 0.634 → 에피소드 리턴 분산 ≈ 0.402. MSE 0.0463과 비교하면
설명분산 ≈ 0.88.
주의 — 이 0.88은 상한 추정치다.
return_std는 에피소드 리턴의 표준편차이고 (ppo.py:1072:episode_return = jnp.sum(transitions.reward, axis=0)),value_loss의 타깃은 GAE 리턴(mb_returns)이다. final candidate는 shaping이 0이고gamma=1.0이라 두 분포가 가깝지만,gae_lambda=0.95의 부트스트랩이 타깃 분산을 축소시키므로 실제 설명분산은 이보다 낮을 수 있다.
용량이 모자란 네트워크는 loss가 높은 지점에서 정체한다. 여기는 낮은 지점에서 정체한다. 잔여 오차 상당 부분은 불완전정보 + 덱 셔플에서 오는 환원 불가능한 분산으로 보인다.
(b) 엔트로피는 계수가 잡아둔 평형이다. entropy_mean이 1.114–1.124에서
500 업데이트 내내 미동도 없다. entropy_coef=0.01과 정책 그래디언트가 이룬
평형이지 용량 한계가 아니다. 파라미터를 늘려도 이 값은 안 변한다.
(불완전정보에서 혼합 전략은 정상이므로 이 자체가 결함은 아니다.)
(c) 죽은 게 아니라 느린 것이다. return_mean은 0.2696 → 0.2916으로
여전히 오르고 있다. approx_kl도 0.042–0.044로 유지되어 정책이 계속 움직인다.
수렴해서 멈춘 상태가 아니다.
4. 실제 천장은 exploitability이고, 이건 용량 문제가 아니다
diminishing-returns-2026-07-05.md 기준,
final 계열을 상대로 새로 학습시킨 exploiter가 여전히 승률 0.539–0.549로 이긴다
(repair_c01_update_500, 3개 프로토콜).
불완전정보 게임에서 이 잔여 착취가능성은 게임이론적 문제다. PPO 셀프플레이는 내쉬로 수렴하지 않고 전략공간을 순환한다. MLP를 키우면 같은 순환 역학 안에서 더 나은 응수를 찾을 뿐, 이 0.54 바닥 자체를 무너뜨리지는 못한다.
판정
- 스케일링이 소폭 이득을 줄 가능성은 있다 (2절: 샘플:파라미터 2,000:1).
- 그러나 용량 부족의 직접 증거는 없다 (3절: 낮은 지점에서 평평한 value_loss, 계수가 고정한 엔트로피).
- 그리고 모델을 실제로 가두고 있는 것은 exploitability이며, 이건 스케일로 풀리지 않는다 (4절).
따라서 우선순위는 낮다. 다만 한 번도 측정한 적이 없다는 사실 때문에 위 판단은 전부 정황증거다. 결론을 확정하려면 단발 A/B가 필요하다 → docs/plans/jax-ppo-model-size-ab.md.
Notes
- 이 작업에서 신규 학습, config 변경, 체크포인트 수정은 없었다.
- 파라미터 수는
ActorCritic(hidden, layers).init()후 리프 크기 합으로 계산했다 (ppo.py:104). - GUI 영향 참고: 현재 CPU 추론 약 50ms/수. 1024×4로 키워도 150–200ms 수준이라
lost-cities-play대국에는 지장이 없다.