From d1c2cf628eb4ce06d305595eeee1068414e09bf7 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EC=A0=95=EC=8B=9C=EC=9B=90?= Date: Wed, 15 Jul 2026 01:27:02 +0900 Subject: [PATCH] Reset envs inside the rollout scan and fix the metrics it breaks rollout_steps was pinned to MAX_STEPS (400) while a round actually runs ~50 plies, and finished envs were only reset between updates. 83% of every rollout was spent stepping already-done envs to produce masked-out zeros. Measured at rollout_steps=400: active steps go 17.3% -> 100%, i.e. 5.8x the learner actions per update for the same compute. Resetting in-scan exposes three things that were previously benign: - compute_gae bootstrapped truncated episodes from zero. That was safe only because every episode used to terminate inside the scan; now episodes cross the boundary, so thread V(s_T) through. - rollout_metrics read final_env and summed rewards along the scan axis, both of which assume one episode per slot. With several episodes per slot that silently produces garbage, so aggregate at done boundaries instead. - league assignments were redrawn only between updates, which would pin a slot to one seat/opponent across every episode in a scan. Redraw them on reset. Also anneal potential shaping against learner actions rather than padded scan steps: the old accounting counted the dead steps, so a 5M-step anneal expired within two updates of 250. Any earlier evidence that shaping does not help was gathered with it effectively off. Co-Authored-By: Claude Opus 4.8 Claude-Session: https://claude.ai/code/session_01XBQKgvBbxbheiTF1AVy1Sh --- docs/plans/lost-cities-classic-3round.md | 226 ++++++++++++++++++ src/lost_cities_jax/ppo.py | 252 ++++++++++++++------- tests/lost_cities_jax/test_ppo_stack.py | 11 +- tests/lost_cities_jax/test_rollout_loop.py | 188 +++++++++++++++ 4 files changed, 587 insertions(+), 90 deletions(-) create mode 100644 docs/plans/lost-cities-classic-3round.md create mode 100644 tests/lost_cities_jax/test_rollout_loop.py 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)