diff --git a/docs/plans/lost-cities-classic-3round.md b/docs/plans/lost-cities-classic-3round.md new file mode 100644 index 0000000..26369e5 --- /dev/null +++ b/docs/plans/lost-cities-classic-3round.md @@ -0,0 +1,226 @@ +# Plan: 클래식 3라운드 로스트시티 에이전트 + +**Status:** Ready to execute (Fable 2차 검토 반영 완료) +**Priority:** High — 현재 스택은 **단판(1라운드)** 게임을 학습하는데, 원작 로스트시티는 +**3라운드 합산**이 승패를 가른다. 즉 지금 에이전트는 애초에 다른 게임을 배우고 있다. +**Background:** Fable 설계 검토 2회 (2026-07-15), 인간 대 AI 기보 +(`data/human-play/game-records.jsonl`, 계속 증가 중) + +## 목표 + +**원작 규칙 그대로의 로스트시티에서 가장 강한 에이전트.** 레거시 코드 호환은 고려하지 않는다. +env/obs/보상/학습 루프 재작성 모두 허용. + +## 확정된 규칙 (룰북 원문 검증 완료) + +코스모스 공식 영문 룰북(691820-02) 및 Board Game Arena 구현체로 교차 확인: + +- 3라운드를 두고 **누적 총점**이 높은 쪽이 매치 승리. +- 1라운드 선공: "가장 나이 많은 사람" → 게임 내적으로는 **임의**. +- **2·3라운드 선공: 누적 점수가 더 많은 쪽** ("The player who has more points begins.") + — 번갈아 가는 것이 아니다. +- **정확한 동점일 때의 선공은 룰북에 없다.** 공식 룰북·BGA 모두 침묵. + → **무작위(동전 던지기)로 정한다.** 결정론적 규칙은 대칭 제로섬 셀프플레이에 + 자리(seat) 비대칭을 주입해 착취 가능한 구멍이 된다. + +점수 계산(`(랭크합 − 20) × (1 + 악수) + (8장 이상 +20)`)과 "방금 버린 카드는 즉시 +회수 불가"(`just_discarded`)는 **현재 엔진이 이미 정확하다** +(`engine.py:244`, `engine.py:124`). + +### 선공이 중요한 이유: 덱 시계 + +라운드는 **덱의 마지막 카드를 뽑는 순간** 끝난다. 카드는 놓은 뒤에 뽑으므로 마지막 +카드를 뽑은 쪽은 그 카드를 쓰지 못한다. 총 턴 수가 홀수면 **선공이 한 장 더 놓는다**. +총 턴 수 = 44 + (버림패 드로우 횟수)이므로 **플레이어가 홀짝을 조작할 수 있다** +(실측 게임당 버림패 드로우 ≈6.6회). 선공권이 누적 점수에 달려 있으므로, 2라운드 +마진에는 **점수를 넘어 3라운드 선공권이라는 추가 가치**가 붙는다. + +## 현재 코드의 결함 (Fable 2회 교차 검증 — 8건 전부 사실 확인) + +1. **env가 1라운드짜리** — `State`에 라운드/누적 점수 없음 (`types.py:46`), + `to_move=0` 하드코딩 (`engine.py:63`). +2. **롤아웃의 87%가 죽은 스텝** — 한 라운드는 실측 평균 50.6수(최대 56)인데 + `rollout_steps = MAX_STEPS = 400` (`ppo.py:66`)이고, 스캔 **안에는 리셋이 없다** + (`reset_done_envs`는 업데이트 사이에서만, `ppo.py:275`). +3. **GAE 부트스트랩 초기값이 0** (`ppo.py:691`). 지금은 모든 에피소드가 스캔 안에서 끝나 + 무해하나, **in-scan 리셋을 켜는 순간 편향**이 된다. +4. **shaping 어닐링 회계 버그** — `update * batch_games * rollout_steps` (`ppo.py:222`)로 + **죽은 스텝까지 세어서**, 5M 어닐링이 250 업데이트 중 2번째에 끝난다. league는 shaping을 + 아예 끈다 (`league.py:592`). + → **"shaping은 불필요하다"는 과거 결론이 있다면 무효다. 켜진 적이 없다.** +5. **obs에 `to_move`가 없다.** 그런데 value loss는 `active` 마스크로 학습되어 (`ppo.py:624`) + **상대 차례 상태에서도** 가치를 맞추라고 요구한다. critic이 차례를 구분할 단서가 + `step_count / 400`뿐 — MLP에게 부동소수점에서 홀짝을 뽑으라는 요구다. +6. **점수차 정규화가 ÷780** (`MAX_ABS_SCORE`, `obs.py:68`). ±50점 차가 ±0.06 → 사실상 + 안 보인다. +7. **셀프플레이에서 상대 자리 결정을 전부 버린다** (`stop_gradient`, `ppo.py:425`) — + 의사결정의 절반을 낭비. +8. **평가가 단판 승률** (`gates.py`) — 3라운드 개편의 성패를 측정할 자가 없다. + +## 핵심 설계 결정 + +### PRNG: 모든 무작위성을 `reset()`에서 미리 뽑는다 (필수) + +라운드 전환 리셔플과 동점 코인플립을 `step()` 안에서 샘플링하면 `step`이 확률적이 되어 +서명이 바뀌고, 모든 rollout/eval/gates 바디에 키를 꿰어야 하며 duplicate 미러링이 꼬인다. + +**`reset()`에서 3라운드 덱 순서 전부(3×60), 동점 코인플립 비트 2개, 1라운드 선공 비트 +1개를 미리 샘플링해 `State`에 저장한다.** `step()`은 결정론을 유지하고 라운드 전환은 +다음 덱 슬라이스로 스위치만 한다. 덤: **미러 매치(같은 3딜 + 자리 교대 + 같은 코인플립 +비트)가 공짜로 따라온다** — Phase 5의 antithetic 페어링이 그대로 성립. + +### 라운드 전환 vs 매치 리셋의 분리 + +- **라운드 전환은 엔진 `step()` 내부**에서 (`done=False` 유지, carry 갱신, 덱 슬라이스 전환) +- **매치 done 리셋은 rollout 바디**에서 (in-scan auto-reset) + +이러면 두 메커니즘이 깨끗이 분리된다. 단 위의 pre-sampled 덱 설계가 전제다. + +- `State`에 **라운드 스텝 카운터와 매치 스텝 카운터를 둘 다** 둔다 (forced-done은 라운드 + 단위, obs 정규화는 라운드 상대 진행도 + round 원-핫). +- 중간 라운드 forced-done(라운드 상한 초과) 시 **자연 종료와 동일하게 라운드 종료 → 전환**. + +### 보상 스케일에 대한 해소 (혼동 방지) + +`sign()`에 가까운 작은 scale은 **틀린 선택**이다 — 150수에 걸쳐 모든 수가 동일한 ±1을 +받아 크레딧 할당이 전부 critic에 떠넘겨지고, 5점 차 패배와 80점 차 패배의 그래디언트가 +같아진다. **scale 30~50의 tanh가 절충점**이다: |마진| ≳ 60에서는 `sign`을 근사하면서 +크레딧 그래디언트를 보존한다. 대신 5점 승과 80점 승을 여전히 구분하므로 순수 승률 +최적화 대비 **마진 쪽으로 약간 왜곡**된다 — 이 잔여 왜곡은 Phase 5의 후기 fine-tune으로 +scale을 낮춰서(예: 50 → 25) 제거한다. + +명시적 risk 항은 **넣지 않는다.** carry + round를 obs에 넣은 종료 보상이 올바른 리스크 +태도(뒤지면 도박, 앞서면 잠금)를 자동으로 유도한다. + +## 체크리스트 + +### Phase 0a — 학습 루프 수정 (최우선, 단판 체제에서 검증 가능) + +- [x] **GAE 부트스트랩 수정**: 잘린 에피소드를 `V(s_final)`로 부트스트랩 (`ppo.py:691`). +- [x] **in-scan auto-reset**: 롤아웃 스캔 안에서 done env 리셋 → 죽은 스텝 제거. + **실측(rollout_steps=400, batch_games=64, 3 업데이트):** main은 활성 17.3% + (learner 액션 2,221/업데이트), 수정 후 **활성 100% (12,842/업데이트) — 동일 연산량에 + 샘플 5.8배.** 활성 비율은 게임 길이에 비례하므로, 학습 초기 정책(~69수)에서 5.8배, + 학습된 정책(~50수)이면 8배에 가까워진다. 3라운드 매치(~150수)가 되면 약 2.7배로 + 줄지만 그때는 스캔 전체가 실제 매치로 채워진다. +- [x] **메트릭 파이프라인 재작업 (필수, 놓치면 조용히 오염됨)**: `rollout_metrics`는 + `final_env`에서 통계를 읽고 `episode_return = jnp.sum(transitions.reward, axis=0)`은 + **env당 에피소드 1개를 가정**한다 (`ppo.py:1069`). in-scan 리셋 후에는 final_env가 + 에피소드 중간이고 여러 에피소드의 보상이 합산된다 → **metrics.jsonl 전체가 쓰레기가 + 된다.** done 경계에서 누적하는 에피소드 단위 집계로 바꿀 것. +- [x] **league assignments를 scan carry로**: `learner_seat`/`opponent_index`/`use_mirror`가 + 업데이트 사이에서만 재샘플링된다 (`ppo.py:568`). in-scan 리셋을 켜면 리셋 시점에 + 스캔 내부에서 재샘플링해야 한다 — 안 하면 자리/상대가 에피소드 간 고정되어 편향. +- [x] **shaping 어닐링 회계 수정**: 패딩된 스캔 스텝이 아니라 **learner 액션 수** 기준 + (`ppo.py:222`). `shaping_coef`는 호스트에서 계산해 jit 함수에 넘기므로 + (`ppo.py:221`), 직전 업데이트의 액션 수를 호스트에서 누적하는 **1-업데이트 지연** + 구조가 된다 (무해). +- [ ] **회귀 게이트**: 단판 체제에서 동일 learner-액션 예산으로 기존 anchor 성적 재현 + ≥ 동등. **이걸 통과해야 다음으로 간다** (성공 기준 0). + +### Phase 1 — 3라운드 env + +- [ ] `State`에 `carry`(누적 점수차), `round_idx`, **라운드/매치 스텝 카운터**, + **pre-sampled 덱 3세트 + 코인플립 비트 + 1라운드 선공 비트** 추가. +- [ ] `step()`: 덱 소진 시 라운드 < 3이면 carry 갱신 → 다음 덱 슬라이스로 전환 → + `round_idx += 1`, `done` 유지. 라운드 == 3이면 `done = True`. +- [ ] **선공 규칙**: 1라운드 pre-sampled 비트. 2·3라운드는 `carry > 0` → 나, `< 0` → 상대, + `== 0` → **pre-sampled 코인플립 비트**. +- [ ] 라운드당 스텝 상한을 400 → **~120**으로 분리. **전역 `MAX_STEPS` 상수를 직접 수정하지 + 말 것** — engine forced-done(`engine.py:224`), obs 정규화(`obs.py:75`), 모든 eval 스캔 + 길이(`ppo.py:773`, `ppo.py:1049`, `gates.py`)에 물려 있다. 매치 레벨 스캔 길이를 따로 둔다. +- [ ] **league 풀 재구축**: 기존 스냅샷 멤버는 obs 변경으로 전부 무효화된다 — Phase 1 산출물로 명시. +- [ ] **테스트** (엔진 재작성에 테스트 없이 들어가면 3라운드 버그를 학습 곡선으로만 발견하게 된다): + 라운드당 총 턴 수 = 44 + 버림패 드로우 (property), `carry` = 라운드 점수 합 불변식, + 선공 규칙 3분기(`>`, `<`, `==`) 각각, 코인플립 통계(≈50/50), PBRS Φ가 라운드 경계에서 + 연속(보드차 → carry 흡수), 미러 매치 대칭성. + +### Phase 0b — 측정 자 (첫 장기 3라운드 학습 **전에** 착지) + +- [ ] **매치 단위 평가**: 3딜 전부 미러링 + 자리 교대 + **동일 코인플립 비트**, + 매치 승률 + Wilson CI, 평균 총 마진. `gates.py`/`league.py`의 단판 기준을 대체. + **`gates.py:884`, `gates.py:961`의 `MAX_STEPS` 길이 단판 스캔 2곳 포팅 포함** — + exploiter 학습·평가 경로 전체가 3라운드로 가야 성공 기준 3을 잴 수 있다. +- [ ] **carry 조건부 프로브**: carry ∈ {−60, −25, −1, +1, +25, +60}을 주입한 3라운드 시작 + 위치에서 승률·행동 변화(원정 개수, 악수 비율, 덱 레이스 비율) 측정. +- [ ] **shuffle bank 포맷 확장**: 덱과 함께 코인플립·선공 비트도 뽑도록. 안 그러면 + `jax.random.PRNGKey(0)` 고정 eval(`ppo.py:1025`)의 재현성이 매치 정의와 얽힌다. +- [ ] **legacy obs 버전 보존**: 구 정책을 3라운드 매치에 투입해 베이스라인을 재려면 구 + obs(454차원)로 추론해야 하는데, `checkpoint_policy_from_params`는 전역 `observation`을 + 호출한다 (`ppo.py:802`) → obs 개편 후 구 체크포인트는 **로드 자체가 실패**한다. + policy 로더가 체크포인트별 obs 버전을 받도록 할 것. **성공 기준 2의 분모가 여기 달렸다.** + +### Phase 2 — 보상 + +- [ ] 종료 보상: 3라운드 끝에서만 `tanh(총_점수차 / scale)`, **scale = 30~50** + (근거는 위 "보상 스케일에 대한 해소" 절). +- [ ] 마진 shaping 부활: `Φ(s) = carry + 현재 보드 점수차`, **작은 계수**(종료 보상 스케일의 + 0.05~0.2). 현재의 계수 1.0 원점수 shaping은 ±1 종료 보상보다 10~30배 크다. + learner 스텝 기준으로 어닐링하되 **후반까지 정확히 0으로 내리지 않는다.** +- [ ] γ=1.0에서 `Φ(s') − Φ(s)`는 **올바른 PBRS 형태다** (γ 누락 아님 — 검토에서 확인). + +### Phase 3 — 관측 + +- [ ] `carry`: **÷75 스칼라 + 구간 원-핫**(약 9구간). 3라운드 정책은 "1점만 더" 임계값 + 근처에서 급격히 꺾여야 한다. 기존 `score_diff`의 ÷780 정규화도 같이 고친다. +- [ ] `round_idx` 원-핫 + 남은 라운드 수. +- [ ] **`to_move` 비트** + "내가 이번 라운드 선공인가" + "현재 홀짝에서 마지막 덱 카드를 + 누가 뽑는가"(덱 시계). +- [ ] 색깔별 **살아있는 점수 3분할**: 내 `col_top` 위로 아직 나올 수 있는 점수를 + (a) **내 손패**, (b) **버림패 더미**(공개돼 있고 회수 가능 — 빠뜨리기 쉬움), + (c) **미공개**(덱 ∪ 상대 은닉 손패)로 나눠 넣는다. 상대에 대해서도 동일 + (상대 `col_top`은 공개). MLP가 7×60 채널에서 뽑아내기 어려운 비선형 집계이고, + 모든 개시/연장/차단 판단을 좌우한다. + +### Phase 4 — 학습 효율 + +- [ ] **전지적 critic (CTDE)**: critic에만 상대 손패 + 덱 구성을 준다. 행동과 무관한 + 정보이므로 정책 그래디언트를 편향시키지 않는다. 딜 운 분산을 정면으로 깎는다. + **가치 경로를 정책 트렁크에서 분리해야 한다** (현재 공유 트렁크, `ppo.py:104`) — + 특권 정보가 정책 로짓으로 새면 안 된다. + → 이 critic은 나중에 **PIMC 탐색의 리프 평가기로 그대로 재활용**된다. +- [ ] **양쪽 자리 학습**: `stop_gradient`된 상대 자리 전이(`ppo.py:425`)도 학습에 쓴다 + (샘플 효율 2배). 같은 게임의 두 자리는 반상관이므로 같은 배치에 두고 advantage + 정규화에 맡긴다. +- [ ] `gae_lambda` 재검토: 0.95는 150수 지평에서 너무 짧다(중반 수 직접 가중치 0.02). + 전지적 critic이 있으면 유지, 없으면 0.97~0.99. **λ=1은 금지**(딜 분산). +- [ ] `entropy_coef` **스윕으로 재결정**. (주의: "0.01이 원점수 shaping 기준으로 잡혔다"는 + 추론은 **틀렸다** — league는 shaping을 끄고 학습했으므로 0.01은 이미 ±1 tanh 체제에서 + 동작해온 값이다. 다만 3라운드에서 보상 빈도가 1/150로 희석되고 작은 shaping이 + 추가되므로 재튜닝 자체는 타당하다. **근거 없이 10배 낮추면 과소탐색으로 직행한다.**) + +### Phase 5 (나중) — 최종 강함 + +- [ ] **후기 fine-tune**: 학습 말미에 종료 보상 scale을 낮춰(50 → 25) 마진 왜곡을 제거하고 + 순수 승률 쪽으로 당긴다. +- [ ] 페어드 antithetic 딜(같은 3딜 + 자리 교대 + 같은 코인플립 비트)을 **학습에** 도입. + 단순 포함이 아니라 **쌍으로 묶어** control variate로 써야 효과가 있다. + (Phase 1의 pre-sampled PRNG 설계 덕에 사실상 공짜.) +- [ ] MMD식 정규화 셀프플레이 (loss에 ~20줄) — 2인 제로섬 근사 내시 보험. +- [ ] **추론 시점 탐색 (PIMC / ISMCTS)** — 최종 강함의 가장 큰 이득. 로스트시티는 블러핑 + 경제가 없는 저기만성 불완전정보 게임이라 결정화 탐색이 잘 맞는다. raw net 대 + net+search 맞대결로 측정. +- [ ] 네트워크 용량 A/B (512×3 → 1024×3 또는 residual) — **파이프라인 변경이 끝난 뒤에.** + +## 범위 밖 (명시적 동결) + +- **웹 클라이언트/ONNX는 별도 계획 전까지 레거시 단판 모델로 동결한다.** obs 개편 즉시 + export 파이프라인(`scripts/export_jax_ppo_onnx.py:70`, manifest `observation_size: 454`), + TS obs 빌더, TS 단판 엔진이 전부 비호환이 된다. 이 선언이 없으면 실행 중 스코프가 + 웹 재작성으로 샌다. +- **Deep CFR 복귀** — 이미 BC 천장을 쳤고, 3라운드는 트리만 키운다. PPO+league를 학습 + 백본으로 유지한다. +- 레거시 호환을 위한 타협. + +## 성공 기준 + +0. **(Phase 0a 게이트)** 루프 수정 후 단판 체제에서 동일 learner-액션 예산으로 기존 anchor + 성적 재현 ≥ 동등. **없으면 auto-reset 버그가 3라운드 결과에 섞여 원인 분리가 불가능해진다.** +1. **carry 프로브에서 행동이 단조롭게 변한다** — carry 6개 수준에 걸쳐 원정 개수/악수 비율의 + **단조 추세**(CI 포함). ("행동이 변한다"는 노이즈로도 통과 가능하므로 단조성으로 정의.) +2. **매치 승률이 구 정책(3라운드에 그대로 투입)을 이긴다** — 페어드 매치 **≥ 1만 쌍**, + 매치 승률 **Wilson 하한 > 0.5** *및* 평균 총 마진 **CI 하한 > 0**. +3. 착취자(exploiter) 승률이 악화되지 않는다. +4. **(anchor 비회귀)** 휴리스틱 anchor 상대 매치 승률·라운드당 마진이 구 정책 대비 악화되지 + 않는다. **기준 2의 구멍을 막는 항목**: carry-blind인 구 정책만 상대로 이기는 것은 + **카드 플레이가 퇴보해도 carry 착취만으로 달성 가능**하다. diff --git a/src/lost_cities_jax/ppo.py b/src/lost_cities_jax/ppo.py index 0f9266c..305a694 100644 --- a/src/lost_cities_jax/ppo.py +++ b/src/lost_cities_jax/ppo.py @@ -77,6 +77,14 @@ class PPOHyperConfig: @dataclass class RewardConfig: + """Terminal reward plus potential-based shaping on the score difference. + + ``potential_shaping_anneal_steps`` counts **learner actions**, not scan + steps. It used to be fed padded scan steps, which made a 5M-step anneal + expire inside the first two updates; any earlier conclusion that shaping + does not help was drawn with shaping effectively off. + """ + terminal_scale: float = 50.0 potential_shaping_initial: float = 1.0 potential_shaping_final: float = 0.0 @@ -134,6 +142,23 @@ class Transition(NamedTuple): play_action: jax.Array +class EpisodeEnd(NamedTuple): + """Per-step episode-completion records emitted by the rollout scan. + + With in-scan auto-reset an env slot hosts several episodes per rollout, so + ``final_env`` is mid-episode and summing rewards over the scan axis mixes + episodes. Every field here is masked by ``done``: read it only where + ``done`` is set. + """ + + done: jax.Array + episode_return: jax.Array + length: jax.Array + opened_colors: jax.Array + positive_expeditions: jax.Array + hit_max_steps: jax.Array + + class EvalBatch(NamedTuple): wins: jax.Array losses: jax.Array @@ -218,14 +243,20 @@ def train( train_iteration = make_train_iteration(cfg, opponent_policy) start = time.perf_counter() + # Anneal against learner actions taken, not scan steps issued: most scan + # steps used to be spent stepping already-done envs, which made the anneal + # finish within the first couple of updates. The count lags by one update + # because it comes back as a device metric. + learner_actions_seen = 0 + for update in range(cfg.run.total_updates): - shaping_coef = shaping_coefficient( - cfg, update * cfg.ppo.batch_games * cfg.ppo.rollout_steps - ) + shaping_coef = shaping_coefficient(cfg, learner_actions_seen) iter_start = time.perf_counter() state, env_state, rng, metrics = train_iteration(state, env_state, rng, shaping_coef) jax.tree_util.tree_leaves(metrics)[0].block_until_ready() row = _metrics_to_row(metrics, update, shaping_coef, cfg) + learner_actions_seen += int(row["learner_actions"]) + row["learner_actions_total"] = learner_actions_seen row["iteration_seconds"] = time.perf_counter() - iter_start row["elapsed_seconds"] = time.perf_counter() - start _append_jsonl(metrics_path, row) @@ -260,8 +291,10 @@ def make_train_iteration(cfg: JaxPPOConfig, opponent_policy): def train_iteration( state: TrainState, env_state: State, rng: jax.Array, shaping_coef: jax.Array ): - rng, rollout_key, update_key, reset_key = jax.random.split(rng, 4) - env_state, transitions, rollout_metrics = rollout_fn( + rng, rollout_key, update_key = jax.random.split(rng, 3) + # The rollout resets finished envs in-scan, so env_state never comes + # back done and needs no reset here. + env_state, transitions, last_value, rollout_stats = rollout_fn( state, env_state, rollout_key, shaping_coef ) advantages, returns = compute_gae( @@ -270,10 +303,10 @@ def make_train_iteration(cfg: JaxPPOConfig, opponent_policy): transitions.done, cfg.ppo.gamma, cfg.ppo.gae_lambda, + last_value, ) state, update_metrics = ppo_update(state, transitions, advantages, returns, update_key, cfg) - env_state = reset_done_envs(env_state, reset_key, cfg.ppo.batch_games) - return state, env_state, rng, {**rollout_metrics, **update_metrics} + return state, env_state, rng, {**rollout_stats, **update_metrics} return train_iteration @@ -299,8 +332,10 @@ def make_league_train_iteration( rng: jax.Array, shaping_coef: jax.Array, ): - rng, rollout_key, update_key, reset_key = jax.random.split(rng, 4) - env_state, transitions, rollout_metrics = rollout_fn( + rng, rollout_key, update_key = jax.random.split(rng, 3) + # The rollout resets finished envs and redraws their league assignments + # in-scan, so neither needs to be refreshed here. + env_state, assignments, transitions, last_value, rollout_stats = rollout_fn( state, env_state, assignments, rollout_key, shaping_coef ) advantages, returns = compute_gae( @@ -309,17 +344,10 @@ def make_league_train_iteration( transitions.done, cfg.ppo.gamma, cfg.ppo.gae_lambda, + last_value, ) state, update_metrics = ppo_update(state, transitions, advantages, returns, update_key, cfg) - env_state, assignments = reset_done_envs_with_assignments( - env_state, - assignments, - reset_key, - cfg.ppo.batch_games, - opponent_probs, - mirror_probability, - ) - return state, env_state, assignments, rng, {**rollout_metrics, **update_metrics} + return state, env_state, assignments, rng, {**rollout_stats, **update_metrics} return train_iteration @@ -328,11 +356,13 @@ def make_rollout_fn(cfg: JaxPPOConfig, opponent_policy): learner = jnp.asarray(cfg.run.learner_seat, dtype=jnp.int32) opponent = jnp.asarray(1 - cfg.run.learner_seat, dtype=jnp.int32) + learners = jnp.full((cfg.ppo.batch_games,), cfg.run.learner_seat, dtype=jnp.int32) + @jax.jit def rollout_fn(state: TrainState, env_state: State, rng: jax.Array, shaping_coef: jax.Array): def body(carry, _): - env, key = carry - key, learner_key, opponent_key = jax.random.split(key, 3) + env, episode_return, key = carry + key, learner_key, opponent_key, reset_key = jax.random.split(key, 4) obs = jax.vmap(observation, in_axes=(0, None))(env, learner) legal = jax.vmap(legal_action_mask)(env) logits, value = state.apply_fn(state.params, obs) @@ -355,7 +385,8 @@ def make_rollout_fn(cfg: JaxPPOConfig, opponent_policy): next_env, _, _ = jax.vmap(step, in_axes=(0, 0))(env, actions) after_diff = batch_score_diff(next_env, learner) terminal_reward = jnp.tanh(after_diff / cfg.reward.terminal_scale) - terminal = active & next_env.done + done = next_env.done + terminal = active & done reward = jnp.where(terminal, terminal_reward, 0.0) reward = reward + shaping_coef * (after_diff - before_diff) reward = jnp.where(active, reward, 0.0) @@ -369,19 +400,27 @@ def make_rollout_fn(cfg: JaxPPOConfig, opponent_policy): log_prob=log_prob, value=value, reward=reward, - done=next_env.done, + done=done, active=active, actor_mask=actor_mask, entropy=entropy, play_action=(place_type == PLAY) & actor_mask, ) - return (next_env, key), transition - (next_env, rng), transitions = jax.lax.scan( - body, (env_state, rng), xs=None, length=cfg.ppo.rollout_steps + episode_return = episode_return + reward + episode = _episode_end(next_env, learners, terminal, episode_return) + next_env = reset_done_envs(next_env, reset_key, cfg.ppo.batch_games) + episode_return = jnp.where(done, 0.0, episode_return) + return (next_env, episode_return, key), (transition, episode) + + init_return = jnp.zeros((cfg.ppo.batch_games,), dtype=jnp.float32) + (next_env, _, rng), (transitions, episodes) = jax.lax.scan( + body, (env_state, init_return, rng), xs=None, length=cfg.ppo.rollout_steps ) - metrics = rollout_metrics(transitions, next_env, learner) - return next_env, transitions, metrics + final_obs = jax.vmap(observation, in_axes=(0, None))(next_env, learner) + _, last_value = state.apply_fn(state.params, final_obs) + metrics = episode_metrics(transitions, episodes) + return next_env, transitions, last_value, metrics return rollout_fn @@ -405,8 +444,8 @@ def make_league_rollout_fn( shaping_coef: jax.Array, ): def body(carry, _): - env, key = carry - key, learner_key, mirror_key, pool_key = jax.random.split(key, 4) + env, assignments, episode_return, key = carry + key, learner_key, mirror_key, pool_key, reset_key = jax.random.split(key, 5) learner = assignments.learner_seat.astype(jnp.int32) opponent = 1 - learner @@ -449,7 +488,8 @@ def make_league_rollout_fn( next_env, _, _ = jax.vmap(step, in_axes=(0, 0))(env, actions) after_diff = batch_score_diff_for_players(next_env, learner) terminal_reward = jnp.tanh(after_diff / cfg.reward.terminal_scale) - terminal = active & next_env.done + done = next_env.done + terminal = active & done reward = jnp.where(terminal, terminal_reward, 0.0) reward = reward + shaping_coef * (after_diff - before_diff) reward = jnp.where(active, reward, 0.0) @@ -463,19 +503,37 @@ def make_league_rollout_fn( log_prob=log_prob, value=value, reward=reward, - done=next_env.done, + done=done, active=active, actor_mask=actor_mask, entropy=entropy, play_action=(place_type == PLAY) & actor_mask, ) - return (next_env, key), transition - (next_env, rng), transitions = jax.lax.scan( - body, (env_state, rng), xs=None, length=cfg.ppo.rollout_steps + episode_return = episode_return + reward + # Scored against the seat that just played the episode out, not the + # freshly drawn one. + episode = _episode_end(next_env, learner, terminal, episode_return) + next_env, assignments = reset_done_envs_with_assignments( + next_env, + assignments, + reset_key, + cfg.ppo.batch_games, + opponent_probs, + mirror_probability, + ) + episode_return = jnp.where(done, 0.0, episode_return) + return (next_env, assignments, episode_return, key), (transition, episode) + + init_return = jnp.zeros((cfg.ppo.batch_games,), dtype=jnp.float32) + (next_env, assignments, _, rng), (transitions, episodes) = jax.lax.scan( + body, (env_state, assignments, init_return, rng), xs=None, length=cfg.ppo.rollout_steps ) - metrics = rollout_metrics_for_players(transitions, next_env, assignments.learner_seat) - return next_env, transitions, metrics + final_learner = assignments.learner_seat.astype(jnp.int32) + final_obs = jax.vmap(observation, in_axes=(0, 0))(next_env, final_learner) + _, last_value = state.apply_fn(state.params, final_obs) + metrics = episode_metrics(transitions, episodes) + return next_env, assignments, transitions, last_value, metrics return rollout_fn @@ -487,11 +545,12 @@ def random_rollout(cfg: JaxPPOConfig) -> dict: env_state = jax.jit(jax.vmap(reset))(jax.random.split(reset_key, cfg.ppo.batch_games)) learner = jnp.asarray(cfg.run.learner_seat, dtype=jnp.int32) opponent = jnp.asarray(1 - cfg.run.learner_seat, dtype=jnp.int32) + learners = jnp.full((cfg.ppo.batch_games,), cfg.run.learner_seat, dtype=jnp.int32) @jax.jit def rollout(env_state: State, rng: jax.Array): def body(carry, _): - env, key = carry + env, episode_return, key = carry key, learner_key, opponent_key = jax.random.split(key, 3) learner_keys = jax.random.split(learner_key, cfg.ppo.batch_games) opponent_keys = jax.random.split(opponent_key, cfg.ppo.batch_games) @@ -508,6 +567,7 @@ def random_rollout(cfg: JaxPPOConfig) -> dict: next_env, _, _ = jax.vmap(step, in_axes=(0, 0))(env, actions) after_diff = batch_score_diff(next_env, learner) reward = jnp.where(active, after_diff - before_diff, 0.0) + terminal = active & next_env.done actor_mask = active & learner_turn place_type = (actions % 12) // 6 dummy_obs = jnp.zeros((cfg.ppo.batch_games, OBS_DIM), dtype=jnp.float32) @@ -525,12 +585,15 @@ def random_rollout(cfg: JaxPPOConfig) -> dict: entropy=jnp.zeros((cfg.ppo.batch_games,), dtype=jnp.float32), play_action=(place_type == PLAY) & actor_mask, ) - return (next_env, key), transition + episode_return = episode_return + reward + episode = _episode_end(next_env, learners, terminal, episode_return) + return (next_env, episode_return, key), (transition, episode) - (next_env, rng), transitions = jax.lax.scan( - body, (env_state, rng), xs=None, length=cfg.ppo.rollout_steps + init_return = jnp.zeros((cfg.ppo.batch_games,), dtype=jnp.float32) + (next_env, _, rng), (transitions, episodes) = jax.lax.scan( + body, (env_state, init_return, rng), xs=None, length=cfg.ppo.rollout_steps ) - return rng, rollout_metrics(transitions, next_env, learner) + return rng, episode_metrics(transitions, episodes) rng, metrics = rollout(env_state, rng) jax.tree_util.tree_leaves(metrics)[0].block_until_ready() @@ -679,7 +742,18 @@ def compute_gae( dones: jax.Array, gamma: float, gae_lambda: float, + last_value: jax.Array | None = None, ) -> tuple[jax.Array, jax.Array]: + """GAE over a rollout that may end mid-episode. + + ``last_value`` is V(s_T) for the state the scan stopped on. Episodes that + are cut by the rollout boundary bootstrap from it; leaving it at zero would + tell the critic every truncated episode is worth nothing. + """ + + if last_value is None: + last_value = jnp.zeros_like(values[-1]) + def body(carry, x): next_value, next_advantage = carry reward, value, done = x @@ -688,7 +762,7 @@ def compute_gae( advantage = delta + gamma * gae_lambda * nonterminal * next_advantage return (value, advantage), advantage - init = (jnp.zeros_like(values[-1]), jnp.zeros_like(values[-1])) + init = (last_value, jnp.zeros_like(values[-1])) _, advantages_rev = jax.lax.scan(body, init, (rewards[::-1], values[::-1], dones[::-1])) advantages = advantages_rev[::-1] return advantages, advantages + values @@ -1066,56 +1140,60 @@ def make_policy_match_batch_fn(learner_policy, opponent_policy): return eval_batch -def rollout_metrics( - transitions: Transition, final_env: State, learner: jax.Array -) -> dict[str, jax.Array]: - episode_return = jnp.sum(transitions.reward, axis=0) +def _episode_end( + next_env: State, + learners: jax.Array, + done: jax.Array, + episode_return: jax.Array, +) -> EpisodeEnd: + """Snapshot the finished-episode stats of ``next_env`` before it is reset.""" + + learners = jnp.broadcast_to(learners.astype(jnp.int32), done.shape) + color_scores = jax.vmap(color_scores_for_player, in_axes=(0, 0))(next_env, learners) + batch_idx = jnp.arange(done.shape[0]) + return EpisodeEnd( + done=done, + episode_return=episode_return, + length=next_env.step_count, + opened_colors=jnp.sum(next_env.col_len[batch_idx, learners, :] > 0, axis=-1), + positive_expeditions=jnp.sum(color_scores > 0, axis=-1), + hit_max_steps=next_env.step_count >= MAX_STEPS, + ) + + +def episode_metrics(transitions: Transition, episodes: EpisodeEnd) -> dict[str, jax.Array]: + """Aggregate over episodes that actually finished inside the rollout. + + Averaging over ``final_env`` instead would sample envs mid-episode, and + summing rewards along the scan axis would add up several episodes per slot. + """ + + done = episodes.done.astype(jnp.float32) + completed = jnp.sum(done) + denom = jnp.maximum(completed, 1.0) actor_count = jnp.sum(transitions.actor_mask) - active_count = jnp.sum(transitions.active) - color_scores = jax.vmap(color_scores_for_player, in_axes=(0, None))(final_env, learner) - opened = jnp.sum(final_env.col_len[:, learner, :] > 0, axis=-1) - positive = jnp.sum(color_scores > 0, axis=-1) + + returns = episodes.episode_return + return_mean = jnp.sum(returns * done) / denom + return_var = jnp.sum(((returns - return_mean) ** 2) * done) / denom + return { - "return_mean": jnp.mean(episode_return), - "return_std": jnp.std(episode_return), - "game_length_mean": jnp.mean(final_env.step_count.astype(jnp.float32)), - "game_length_max": jnp.max(final_env.step_count), - "max_steps_rate": jnp.mean(final_env.step_count >= MAX_STEPS), + "return_mean": return_mean, + "return_std": jnp.sqrt(jnp.maximum(return_var, 0.0)), + "game_length_mean": jnp.sum(episodes.length.astype(jnp.float32) * done) / denom, + "game_length_max": jnp.max(jnp.where(episodes.done, episodes.length, 0)), + "max_steps_rate": jnp.sum(episodes.hit_max_steps.astype(jnp.float32) * done) / denom, "play_action_rate": jnp.sum(transitions.play_action) / jnp.maximum(actor_count, 1), - "opened_colors_mean": jnp.mean(opened.astype(jnp.float32)), - "positive_expeditions_mean": jnp.mean(positive.astype(jnp.float32)), + "opened_colors_mean": jnp.sum(episodes.opened_colors.astype(jnp.float32) * done) / denom, + "positive_expeditions_mean": ( + jnp.sum(episodes.positive_expeditions.astype(jnp.float32) * done) / denom + ), "entropy_mean": masked_mean( transitions.entropy, transitions.actor_mask.astype(jnp.float32) ), - "active_steps": active_count, - "learner_actions": actor_count, - } - - -def rollout_metrics_for_players( - transitions: Transition, final_env: State, learners: jax.Array -) -> dict[str, jax.Array]: - episode_return = jnp.sum(transitions.reward, axis=0) - actor_count = jnp.sum(transitions.actor_mask) - active_count = jnp.sum(transitions.active) - color_scores = jax.vmap(color_scores_for_player, in_axes=(0, 0))(final_env, learners) - batch_idx = jnp.arange(learners.shape[0]) - opened = jnp.sum(final_env.col_len[batch_idx, learners.astype(jnp.int32), :] > 0, axis=-1) - positive = jnp.sum(color_scores > 0, axis=-1) - return { - "return_mean": jnp.mean(episode_return), - "return_std": jnp.std(episode_return), - "game_length_mean": jnp.mean(final_env.step_count.astype(jnp.float32)), - "game_length_max": jnp.max(final_env.step_count), - "max_steps_rate": jnp.mean(final_env.step_count >= MAX_STEPS), - "play_action_rate": jnp.sum(transitions.play_action) / jnp.maximum(actor_count, 1), - "opened_colors_mean": jnp.mean(opened.astype(jnp.float32)), - "positive_expeditions_mean": jnp.mean(positive.astype(jnp.float32)), - "entropy_mean": masked_mean( - transitions.entropy, transitions.actor_mask.astype(jnp.float32) - ), - "active_steps": active_count, + "active_steps": jnp.sum(transitions.active), "learner_actions": actor_count, + "episodes_completed": completed, } @@ -1179,11 +1257,11 @@ def _flatten_transitions(transitions: Transition) -> Transition: return jax.tree_util.tree_map(lambda x: x.reshape((-1, *x.shape[2:])), transitions) -def shaping_coefficient(cfg: JaxPPOConfig, env_steps: int) -> float: +def shaping_coefficient(cfg: JaxPPOConfig, learner_actions: int) -> float: reward_cfg = cfg.reward if reward_cfg.potential_shaping_anneal_steps <= 0: return reward_cfg.potential_shaping_final - progress = min(env_steps / reward_cfg.potential_shaping_anneal_steps, 1.0) + progress = min(learner_actions / reward_cfg.potential_shaping_anneal_steps, 1.0) return reward_cfg.potential_shaping_initial + progress * ( reward_cfg.potential_shaping_final - reward_cfg.potential_shaping_initial ) diff --git a/tests/lost_cities_jax/test_ppo_stack.py b/tests/lost_cities_jax/test_ppo_stack.py index 3aef0e7..35dcfba 100644 --- a/tests/lost_cities_jax/test_ppo_stack.py +++ b/tests/lost_cities_jax/test_ppo_stack.py @@ -61,7 +61,9 @@ def tiny_config(tmp_path) -> JaxPPOConfig: ), opponent=OpponentConfig(name="discard_only"), network=NetworkConfig(hidden_size=32, num_layers=1), - ppo=PPOHyperConfig(batch_games=8, rollout_steps=16, epochs=1, minibatches=2), + # A round runs at least 44 plies (one per deck draw), so a shorter + # rollout would finish no episodes and leave the episode metrics empty. + ppo=PPOHyperConfig(batch_games=8, rollout_steps=80, epochs=1, minibatches=2), ) @@ -305,9 +307,12 @@ def test_gate3_checkpoint_duplicate_self_mirror_score_diff_is_zero(): def test_random_rollout_smoke(tmp_path): row = random_rollout(tiny_config(tmp_path)) - assert row["env_steps"] == 8 * 16 + assert row["env_steps"] == 8 * 80 assert 0.0 <= row["play_action_rate"] <= 1.0 - assert row["game_length_mean"] > 0.0 + # Averaged over episodes that actually finished, so it must clear the + # 44-ply floor rather than report a mid-episode step count. + assert row["episodes_completed"] > 0 + assert row["game_length_mean"] >= 44.0 def test_train_checkpoint_and_eval_smoke(tmp_path): diff --git a/tests/lost_cities_jax/test_rollout_loop.py b/tests/lost_cities_jax/test_rollout_loop.py new file mode 100644 index 0000000..9e519fd --- /dev/null +++ b/tests/lost_cities_jax/test_rollout_loop.py @@ -0,0 +1,188 @@ +"""Phase 0a: GAE bootstrap, in-scan auto-reset, and episode-boundary metrics.""" + +import jax +import jax.numpy as jnp +import pytest + +from lost_cities_jax.engine import reset +from lost_cities_jax.opponents import policy_by_name +from lost_cities_jax.ppo import ( + EpisodeEnd, + JaxPPOConfig, + Transition, + compute_gae, + create_train_state, + episode_metrics, + make_rollout_fn, + shaping_coefficient, +) + + +def _cfg(**overrides) -> JaxPPOConfig: + cfg = JaxPPOConfig() + cfg.ppo.batch_games = 8 + cfg.ppo.rollout_steps = 120 + cfg.network.hidden_size = 16 + cfg.network.num_layers = 1 + for key, value in overrides.items(): + setattr(cfg.run, key, value) + return cfg + + +# --- GAE bootstrap --------------------------------------------------------- + + +def test_gae_bootstraps_truncated_episode_from_last_value(): + """A rollout that ends mid-episode must carry V(s_T), not zero.""" + rewards = jnp.zeros((3, 1)) + values = jnp.zeros((3, 1)) + dones = jnp.zeros((3, 1), dtype=bool) # never terminates inside the window + + _, returns_zero = compute_gae(rewards, values, dones, gamma=1.0, gae_lambda=1.0) + _, returns_boot = compute_gae( + rewards, values, dones, gamma=1.0, gae_lambda=1.0, last_value=jnp.array([5.0]) + ) + + # Without a bootstrap the truncated episode looks worthless. + assert jnp.allclose(returns_zero, 0.0) + # With it, every step inherits the tail value. + assert jnp.allclose(returns_boot, 5.0) + + +def test_gae_ignores_last_value_when_episode_terminates(): + """A terminal step must not bootstrap past the episode boundary.""" + rewards = jnp.array([[0.0], [1.0]]) + values = jnp.zeros((2, 1)) + dones = jnp.array([[False], [True]]) + + _, returns = compute_gae( + rewards, values, dones, gamma=1.0, gae_lambda=1.0, last_value=jnp.array([99.0]) + ) + # Terminal step sees only its own reward; the step before it sees that too. + assert jnp.allclose(returns, jnp.array([[1.0], [1.0]])) + + +def test_gae_done_flag_cuts_credit_between_episodes(): + """With in-scan resets two episodes share a slot; credit must not leak.""" + rewards = jnp.array([[1.0], [7.0]]) + values = jnp.zeros((2, 1)) + dones = jnp.array([[True], [False]]) # first step ends episode A + + _, returns = compute_gae(rewards, values, dones, gamma=1.0, gae_lambda=1.0) + # Episode A keeps its own reward; episode B's +7 must not flow backwards. + assert jnp.allclose(returns[0], 1.0) + assert jnp.allclose(returns[1], 7.0) + + +# --- episode-boundary metrics --------------------------------------------- + + +def _episodes(done, episode_return, length) -> EpisodeEnd: + shape = jnp.asarray(done).shape + return EpisodeEnd( + done=jnp.asarray(done), + episode_return=jnp.asarray(episode_return, dtype=jnp.float32), + length=jnp.asarray(length, dtype=jnp.int32), + opened_colors=jnp.zeros(shape, dtype=jnp.int32), + positive_expeditions=jnp.zeros(shape, dtype=jnp.int32), + hit_max_steps=jnp.zeros(shape, dtype=bool), + ) + + +def _transitions(steps: int, games: int) -> Transition: + zeros = jnp.zeros((steps, games), dtype=jnp.float32) + false = jnp.zeros((steps, games), dtype=bool) + return Transition( + obs=jnp.zeros((steps, games, 1)), + legal_mask=jnp.zeros((steps, games, 1), dtype=bool), + action=jnp.zeros((steps, games), dtype=jnp.int32), + log_prob=zeros, + value=zeros, + reward=zeros, + done=false, + active=~false, + actor_mask=~false, + entropy=zeros, + play_action=false, + ) + + +def test_episode_metrics_only_counts_finished_episodes(): + """Unfinished episodes carry a partial return; they must not be averaged in.""" + # Slot 0 finishes twice (returns 10 and 20); slot 1 never finishes. + done = jnp.array([[True, False], [False, False], [True, False]]) + episode_return = jnp.array([[10.0, 3.0], [0.0, 6.0], [20.0, 9.0]]) + length = jnp.array([[50, 1], [0, 2], [60, 3]]) + + metrics = episode_metrics(_transitions(3, 2), _episodes(done, episode_return, length)) + + assert float(metrics["episodes_completed"]) == 2.0 + assert float(metrics["return_mean"]) == pytest.approx(15.0) + assert float(metrics["game_length_mean"]) == pytest.approx(55.0) + assert int(metrics["game_length_max"]) == 60 + + +def test_episode_metrics_survives_a_rollout_with_no_completions(): + done = jnp.zeros((2, 2), dtype=bool) + metrics = episode_metrics( + _transitions(2, 2), _episodes(done, jnp.zeros((2, 2)), jnp.zeros((2, 2))) + ) + assert float(metrics["episodes_completed"]) == 0.0 + assert float(metrics["return_mean"]) == 0.0 # guarded denominator, not a NaN + + +# --- in-scan auto-reset ---------------------------------------------------- + + +def test_rollout_never_returns_a_done_env_and_completes_many_episodes(): + """The whole point of in-scan reset: no dead steps, several games per slot.""" + cfg = _cfg() + rng = jax.random.PRNGKey(0) + rng, init_key, reset_key, roll_key = jax.random.split(rng, 4) + state = create_train_state(cfg, init_key) + env_state = jax.jit(jax.vmap(reset))(jax.random.split(reset_key, cfg.ppo.batch_games)) + + rollout_fn = make_rollout_fn(cfg, policy_by_name("discard_only")) + next_env, transitions, last_value, metrics = rollout_fn( + state, env_state, roll_key, jnp.asarray(0.0, dtype=jnp.float32) + ) + + # Every slot is mid-episode, never parked on a finished game. + assert not bool(jnp.any(next_env.done)) + # 120 steps at ~50 plies a game means each of the 8 slots finished at least one. + assert float(metrics["episodes_completed"]) >= cfg.ppo.batch_games + # No step is wasted stepping an already-done env. + assert bool(jnp.all(transitions.active)) + assert int(metrics["active_steps"]) == cfg.ppo.rollout_steps * cfg.ppo.batch_games + assert last_value.shape == (cfg.ppo.batch_games,) + + +def test_rollout_game_length_matches_the_rules_clock(): + """A round is 44 deck draws plus one turn per discard-pile draw.""" + cfg = _cfg() + rng = jax.random.PRNGKey(1) + rng, init_key, reset_key, roll_key = jax.random.split(rng, 4) + state = create_train_state(cfg, init_key) + env_state = jax.jit(jax.vmap(reset))(jax.random.split(reset_key, cfg.ppo.batch_games)) + + rollout_fn = make_rollout_fn(cfg, policy_by_name("discard_only")) + _, _, _, metrics = rollout_fn(state, env_state, roll_key, jnp.asarray(0.0, dtype=jnp.float32)) + + # 44 deck draws is the floor; discard-pile draws only ever extend a round. + assert float(metrics["game_length_mean"]) >= 44.0 + assert float(metrics["max_steps_rate"]) == 0.0 + + +# --- shaping anneal accounting -------------------------------------------- + + +def test_shaping_anneal_tracks_learner_actions_not_padded_steps(): + cfg = _cfg() + cfg.reward.potential_shaping_initial = 1.0 + cfg.reward.potential_shaping_final = 0.0 + cfg.reward.potential_shaping_anneal_steps = 1_000 + + assert shaping_coefficient(cfg, 0) == pytest.approx(1.0) + assert shaping_coefficient(cfg, 500) == pytest.approx(0.5) + assert shaping_coefficient(cfg, 1_000) == pytest.approx(0.0) + assert shaping_coefficient(cfg, 10_000) == pytest.approx(0.0)