Files
coorl-lost-cities/docs/reports/jax-ppo-model-capacity-2026-07-12.md
T
2026-07-14 20:09:03 +09:00

120 lines
6.1 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
# 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](../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`, `legacy/deep-cfr/configs/` 기준). 특히
[docs/plans/model_size_experiment.md](../plans/archive/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 |
| --- | ---: | ---: | ---: | ---: | ---: |
| 050 | 0.0468 | 1.1159 | 0.0433 | 0.2696 | 0.6299 |
| 50100 | 0.0462 | 1.1167 | 0.0438 | 0.2799 | 0.6275 |
| 100200 | 0.0461 | 1.1241 | 0.0441 | 0.2822 | 0.6316 |
| 200300 | 0.0459 | 1.1143 | 0.0437 | 0.2857 | 0.6306 |
| 300400 | 0.0460 | 1.1239 | 0.0424 | 0.2900 | 0.6323 |
| 400500 | 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.1141.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](diminishing-returns-2026-07-05.md) 기준,
final 계열을 상대로 새로 학습시킨 exploiter가 **여전히 승률 0.5390.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](../plans/jax-ppo-model-size-ab.md).
## Notes
- 이 작업에서 신규 학습, config 변경, 체크포인트 수정은 없었다.
- 파라미터 수는 `ActorCritic(hidden, layers).init()` 후 리프 크기 합으로 계산했다
(`ppo.py:104`).
- GUI 영향 참고: 현재 CPU 추론 약 50ms/수. 1024×4로 키워도 150200ms 수준이라
`lost-cities-play` 대국에는 지장이 없다.