Pay the match total densely and drop the tanh terminal reward
The user asked why we were not simply maximising the three-round total, and they were right. Each ply now pays the points by which it moved the running match difference; at gamma=1 that telescopes to the final total, so the objective is exactly the rulebook's, handed out every ply instead of once 150 plies later. Three measurements killed the tanh design: - Rounds are independent (corr(m1,m2)=0.004, corr(m1+m2,m3)=0.05), so a reward linear in the total decomposes the match into three independent rounds and carry enters the objective nowhere. The only coupling, the start-player rule, is worth +0.73 +/- 0.84 points -- indistinguishable from zero. - Risk attitude, the one thing tanh buys, is worthless here. A policy made to gamble when it trails by 20 entering round three *loses* to a greedy clone over 6144 duplicate matches (0.482); gambling only at -40 breaks even (0.498). A marginal wager buys about +1.7 sigma for -2 to -3 expected points. Ceiling on the whole carry-conditioning idea: under one win-rate point. - Head to head over 10,000 duplicate matches at equal compute, the linear reward *beats* tanh(total/12): 0.5859 (CI 0.576-0.596), +20.3 points. Dropping it is not merely free, it is better -- not because of risk, but because tanh hands a ~150-ply match one saturated +/-1 and leaves all credit assignment to the critic. The flat carry probe was not exploration collapse: sampled play still opens 5.00 expeditions, entropy settles at 1.36 nats (3.9 effective actions), and the critic reads carry cleanly (round-three values run -0.87 to +0.86, monotone). The signal was there; there was nothing to buy with it. Criterion 1 (a monotone carry response) comes off the gate accordingly -- the optimal response barely exists in this game. carry stays in the observation: it costs nothing and the start-player rule keys off it. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01XBQKgvBbxbheiTF1AVy1Sh
This commit is contained in:
@@ -355,3 +355,68 @@ carry 프로브(3라운드 시작 시 carry 주입)에서:
|
|||||||
- [ ] **"항상 5색 개시"가 정상인지 검증**: 인간 기보에서 AI는 4.81, 인간은 4.19를 열었다.
|
- [ ] **"항상 5색 개시"가 정상인지 검증**: 인간 기보에서 AI는 4.81, 인간은 4.19를 열었다.
|
||||||
5.00 고정은 의심스럽다. 셀프플레이 균형인지, 탐색 붕괴인지 (엔트로피 3.5 → 1.03).
|
5.00 고정은 의심스럽다. 셀프플레이 균형인지, 탐색 붕괴인지 (엔트로피 3.5 → 1.03).
|
||||||
- [ ] scale 스케줄(50 → 12 후기 fine-tune)이 처음부터 12로 학습하는 것보다 나은지 A/B.
|
- [ ] scale 스케줄(50 → 12 후기 fine-tune)이 처음부터 12로 학습하는 것보다 나은지 A/B.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 판정: 보상은 **선형 총점**이다 (2026-07-15, Fable 3차 검토 + 실측)
|
||||||
|
|
||||||
|
사용자 제안("그냥 3판 총점이 크기만 하면 되는 것 아니냐")이 **맞았다. tanh 종료 보상은
|
||||||
|
오버엔지니어링이었다.**
|
||||||
|
|
||||||
|
### 근거 1 — 분해 논증 (실측)
|
||||||
|
|
||||||
|
보상이 총점에 선형이면 라운드가 독립이므로 3라운드 게임이 3개의 독립 단판으로 분해된다.
|
||||||
|
셀프플레이 2,048매치 실측: `corr(m1, m2) = 0.004`, `corr(m1+m2, m3) = 0.05`.
|
||||||
|
라운드를 잇는 유일한 고리인 **선공 어드밴티지는 +0.73 ± 0.84점** — 0과 구분 불가.
|
||||||
|
|
||||||
|
### 근거 2 — 리스크 태도는 값어치가 없다 (실측)
|
||||||
|
|
||||||
|
Fable이 리스크 태도를 직접 구현해 greedy 클론과 duplicate 6,144판 맞대결:
|
||||||
|
|
||||||
|
| 도박 정책 | 매치 승률 |
|
||||||
|
|---|---|
|
||||||
|
| 3R에서 20점 이상 뒤지면 온도 샘플링 | **0.482** (진다) |
|
||||||
|
| 40점 이상 뒤질 때만 | 0.498 (본전) |
|
||||||
|
|
||||||
|
**일부러 도박을 시켜도 지거나 본전.** 로스트시티의 분산 레버(한계 악수 ≈ Δσ +1.7에 마진
|
||||||
|
−2~3점)가 근본적으로 약해서, 분산을 사는 비용이 볼록성 이득을 먹는다.
|
||||||
|
carry 조건부 플레이의 가치 상한: **1승점 미만.**
|
||||||
|
|
||||||
|
### 근거 3 — 맞대결에서 단순한 쪽이 **이겼다**
|
||||||
|
|
||||||
|
duplicate 10,000판 (같은 3딜 + 같은 코인 + 자리 교대), 동일 컴퓨트 300 업데이트:
|
||||||
|
|
||||||
|
| A | B | A 승률 | A 평균 총점차 |
|
||||||
|
|---|---|---|---|
|
||||||
|
| **선형 총점** | tanh(총점/12) | **0.5859** (CI 0.576–0.596) | **+20.3점** |
|
||||||
|
|
||||||
|
버려도 손해가 없는 게 아니라 **버리니 더 강해졌다.** 이유는 리스크가 아니라 **신호 밀도**다:
|
||||||
|
tanh는 ~150수 매치에 포화된 ±1 하나를 주고, 선형은 매 수마다 그 수가 총점차를 움직인 만큼을
|
||||||
|
준다. γ=1에서 후자의 합이 정확히 최종 총점차로 telescoping되므로 **목적함수는 동일한데 크레딧
|
||||||
|
할당만 150배 조밀**하다.
|
||||||
|
|
||||||
|
### 근거 4 — 탐색 붕괴가 아니었다
|
||||||
|
|
||||||
|
프로브가 평평했던 이유는 탐색 붕괴가 아니다. 샘플링 프로브에서도 원정 5.00 ± 0.05,
|
||||||
|
엔트로피 1.36나트(유효 행동 3.9개)로 정상 수렴. 그리고 **critic은 carry를 완벽히 읽고 있었다**
|
||||||
|
(3R 시작 가치 −0.87 → +0.86 단조). 신호는 있었고, **정책이 그걸로 살 수 있는 물건이
|
||||||
|
없었을 뿐**이다.
|
||||||
|
|
||||||
|
### 최종 설계
|
||||||
|
|
||||||
|
```python
|
||||||
|
# match_ppo.py 롤아웃 바디
|
||||||
|
reward0 = jnp.where(active, (after - before) / cfg.reward.terminal_scale, 0.0)
|
||||||
|
```
|
||||||
|
|
||||||
|
- **성공 기준 1(carry 프로브 단조성)은 게이트에서 제거한다.** 최적 반응 자체가 이 게임에서
|
||||||
|
거의 존재하지 않는다는 것이 측정 결과다. 기준 2·4(맞대결 + anchor 비회귀)가 옳은 자다.
|
||||||
|
- `carry`는 obs에 **남긴다** (비용 0, 선공 규칙이 키로 쓰는 상태, critic이 잘 읽음).
|
||||||
|
|
||||||
|
### 목적함수와 무관하게 살아남은 것 (전부 순이득)
|
||||||
|
|
||||||
|
in-scan auto-reset + GAE 절단 부트스트랩(샘플 5.8배), 메트릭 재작업, **비대칭 CTDE critic**,
|
||||||
|
**양쪽 좌석 학습**, **duplicate 매치 평가**, pre-sampled PRNG/미러 설계, 그리고 match_obs의
|
||||||
|
단판 결함 수정분(`to_move`, 덱 시계, ÷780 → ÷75, 살아있는 점수 3분할).
|
||||||
|
|
||||||
|
**죽은 것은 tanh 종료 보상과 carry 구간 원-핫뿐이다.**
|
||||||
|
|||||||
@@ -145,6 +145,61 @@ def match_evaluate(cfg: JaxPPOConfig, params, *, matches: int = 2000, seed: int
|
|||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def head_to_head(
|
||||||
|
cfg: JaxPPOConfig, params_a, params_b, *, matches: int = 4096, seed: int = 20260717
|
||||||
|
) -> dict:
|
||||||
|
"""Duplicate matches between two policies: same deals, same coins, both seats.
|
||||||
|
|
||||||
|
Every deal is played twice with the seats swapped, so deal luck cancels and
|
||||||
|
what is left is the difference between the two policies.
|
||||||
|
"""
|
||||||
|
|
||||||
|
decks, coins = match_bank(seed, matches)
|
||||||
|
model = MatchActorCritic(cfg.network.hidden_size, cfg.network.num_layers)
|
||||||
|
|
||||||
|
@jax.jit
|
||||||
|
def run(env: MatchState, a_seat: jax.Array):
|
||||||
|
def body(carry, _):
|
||||||
|
env, _unused = carry
|
||||||
|
to_move = env.round.to_move.astype(jnp.int32)
|
||||||
|
obs = jax.vmap(match_observation)(env, to_move)
|
||||||
|
critic = jax.vmap(match_critic_observation)(env, to_move)
|
||||||
|
mask = jax.vmap(match_legal_action_mask)(env)
|
||||||
|
|
||||||
|
logits_a, _ = model.apply(params_a, obs, critic)
|
||||||
|
logits_b, _ = model.apply(params_b, obs, critic)
|
||||||
|
act_a = jnp.argmax(mask_logits(logits_a, mask), axis=-1).astype(jnp.int32)
|
||||||
|
act_b = jnp.argmax(mask_logits(logits_b, mask), axis=-1).astype(jnp.int32)
|
||||||
|
|
||||||
|
action = jnp.where(to_move == a_seat, act_a, act_b)
|
||||||
|
env, _, _ = jax.vmap(match_step, in_axes=(0, 0))(env, action)
|
||||||
|
return (env, _unused), None
|
||||||
|
|
||||||
|
(env, _), _ = jax.lax.scan(body, (env, jnp.int32(0)), xs=None, length=MATCH_SCAN_STEPS)
|
||||||
|
return env
|
||||||
|
|
||||||
|
leads = []
|
||||||
|
for seat in (0, 1):
|
||||||
|
env = jax.vmap(match_reset_from)(decks, coins)
|
||||||
|
final = run(env, jnp.full((matches,), seat, dtype=jnp.int32))
|
||||||
|
totals = np.asarray(jax.vmap(match_score)(final))
|
||||||
|
leads.append(totals[:, seat] - totals[:, 1 - seat])
|
||||||
|
|
||||||
|
lead = np.concatenate(leads)
|
||||||
|
games = float(lead.size)
|
||||||
|
wins = float((lead > 0).sum())
|
||||||
|
ties = float((lead == 0).sum())
|
||||||
|
low, high = _wilson(wins, games)
|
||||||
|
return {
|
||||||
|
"matches": games,
|
||||||
|
"a_win_rate": wins / games,
|
||||||
|
"wilson_low": low,
|
||||||
|
"wilson_high": high,
|
||||||
|
"ties": ties,
|
||||||
|
"a_mean_lead": float(lead.mean()),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
def carry_probe(
|
def carry_probe(
|
||||||
cfg: JaxPPOConfig,
|
cfg: JaxPPOConfig,
|
||||||
params,
|
params,
|
||||||
|
|||||||
@@ -16,13 +16,24 @@ with its own seat-relative view, value and reward; the actor mask keeps the
|
|||||||
policy loss on the seat that actually moved, while the critic learns from both.
|
policy loss on the seat that actually moved, while the critic learns from both.
|
||||||
Folding the seat axis into the batch keeps each seat's GAE chain independent.
|
Folding the seat axis into the batch keeps each seat's GAE chain independent.
|
||||||
|
|
||||||
**Reward is the match, not the round.** Rounds one and two pay nothing, they only
|
**Reward is the match total, paid densely.** Each ply pays the points by which it
|
||||||
bank into ``carry``. The terminal reward is ``tanh(total_diff / terminal_scale)``
|
moved the running match difference; at gamma=1 that telescopes to the final total
|
||||||
at the end of round three. Driving that scale toward zero would make it
|
difference, so the objective is exactly the rulebook's -- score more across the
|
||||||
``sign()`` -- the true objective, but a poor signal, since every ply of a
|
three rounds.
|
||||||
~160-ply match would then carry the same +/-1 and all credit assignment would
|
|
||||||
fall to the critic. Potential shaping on the running total covers the gap early
|
This used to be ``tanh(total / terminal_scale)`` at the end of round three, to
|
||||||
and anneals away.
|
optimise P(win match) rather than expected total. That was over-engineered. The
|
||||||
|
two differ only in risk attitude, and risk attitude is worth almost nothing here:
|
||||||
|
a policy told to gamble when it trails by 20 going into round three *loses* to a
|
||||||
|
greedy clone over 6144 duplicate matches (0.482), and one that gambles only when
|
||||||
|
it trails by 40 breaks even (0.498). Lost Cities' variance levers -- a marginal
|
||||||
|
wager, a marginal expedition -- cost more expected margin than the convexity they
|
||||||
|
buy. Measured ceiling on the whole carry-conditioning idea: under one win-rate
|
||||||
|
point. The dense linear reward also gives credit assignment every ply instead of
|
||||||
|
one bounded number per ~160.
|
||||||
|
|
||||||
|
``carry`` stays in the observation: it costs nothing, it is the state the
|
||||||
|
start-player rule keys off, and the critic reads it cleanly.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
@@ -208,9 +219,11 @@ def make_match_rollout_fn(cfg: JaxPPOConfig):
|
|||||||
|
|
||||||
done = next_env.done
|
done = next_env.done
|
||||||
terminal = active & done
|
terminal = active & done
|
||||||
reward0 = jnp.where(terminal, jnp.tanh(after / cfg.reward.terminal_scale), 0.0)
|
# Every ply pays the points by which it moved the match total. At
|
||||||
reward0 = reward0 + shaping_coef * (after - before) / cfg.reward.terminal_scale
|
# gamma=1 that telescopes to the final total difference, so the
|
||||||
reward0 = jnp.where(active, reward0, 0.0)
|
# objective is exactly "score more across the three rounds" -- just
|
||||||
|
# handed out densely instead of once, 150 plies later.
|
||||||
|
reward0 = jnp.where(active, (after - before) / cfg.reward.terminal_scale, 0.0)
|
||||||
# Zero-sum: seat 1 sees exactly the negation.
|
# Zero-sum: seat 1 sees exactly the negation.
|
||||||
reward = jnp.stack([reward0, -reward0], axis=0)
|
reward = jnp.stack([reward0, -reward0], axis=0)
|
||||||
|
|
||||||
|
|||||||
@@ -74,7 +74,7 @@ def test_privileged_input_cannot_move_the_policy_logits():
|
|||||||
# --- the match reward ------------------------------------------------------
|
# --- the match reward ------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def test_reward_is_zero_sum_and_only_paid_at_the_end_of_the_match():
|
def test_reward_is_zero_sum():
|
||||||
cfg = _cfg()
|
cfg = _cfg()
|
||||||
state = create_match_train_state(cfg, jax.random.PRNGKey(3))
|
state = create_match_train_state(cfg, jax.random.PRNGKey(3))
|
||||||
env = jax.jit(jax.vmap(match_reset))(
|
env = jax.jit(jax.vmap(match_reset))(
|
||||||
@@ -82,37 +82,37 @@ def test_reward_is_zero_sum_and_only_paid_at_the_end_of_the_match():
|
|||||||
)
|
)
|
||||||
|
|
||||||
rollout = make_match_rollout_fn(cfg)
|
rollout = make_match_rollout_fn(cfg)
|
||||||
# Shaping off, so any non-zero reward must be a terminal one.
|
_, transitions, _, metrics = rollout(state, env, jax.random.PRNGKey(5), jnp.asarray(1.0))
|
||||||
_, transitions, _, metrics = rollout(state, env, jax.random.PRNGKey(5), jnp.asarray(0.0))
|
|
||||||
|
|
||||||
reward = np.asarray(transitions.reward) # (T, 2 * batch)
|
reward = np.asarray(transitions.reward) # (T, 2 * batch)
|
||||||
seat0, seat1 = reward[:, : cfg.ppo.batch_games], reward[:, cfg.ppo.batch_games :]
|
seat0, seat1 = reward[:, : cfg.ppo.batch_games], reward[:, cfg.ppo.batch_games :]
|
||||||
assert np.allclose(seat0, -seat1)
|
assert np.allclose(seat0, -seat1)
|
||||||
|
|
||||||
done = np.asarray(transitions.done)[:, : cfg.ppo.batch_games]
|
|
||||||
paid = seat0 != 0.0
|
|
||||||
assert np.array_equal(paid, done & paid) # never paid on a non-terminal ply
|
|
||||||
assert float(metrics["matches_completed"]) > 0
|
assert float(metrics["matches_completed"]) > 0
|
||||||
assert np.abs(seat0[paid]).max() <= 1.0 # tanh-bounded
|
|
||||||
|
|
||||||
|
|
||||||
def test_shaping_tracks_the_running_match_total():
|
def test_the_dense_reward_telescopes_to_the_match_total():
|
||||||
"""Phi is carry + board diff, so shaping must follow the total, not the round."""
|
"""The whole justification for paying every ply: at gamma=1 the sum is the total.
|
||||||
|
|
||||||
|
That is what makes "score more across three rounds" the objective, rather than
|
||||||
|
some shaped proxy for it.
|
||||||
|
"""
|
||||||
match = match_reset(jax.random.PRNGKey(6))
|
match = match_reset(jax.random.PRNGKey(6))
|
||||||
key = jax.random.PRNGKey(7)
|
key = jax.random.PRNGKey(7)
|
||||||
|
scale = 50.0
|
||||||
|
|
||||||
prev = np.asarray(match_score(match))
|
paid = 0.0
|
||||||
for _ in range(200):
|
while not bool(match.done):
|
||||||
if bool(match.done):
|
|
||||||
break
|
|
||||||
key, step_key = jax.random.split(key)
|
key, step_key = jax.random.split(key)
|
||||||
action = random_legal_action(match.round, match.round.to_move, step_key)
|
action = random_legal_action(match.round, match.round.to_move, step_key)
|
||||||
|
before = np.asarray(match_score(match))
|
||||||
match, _, _ = match_step(match, action)
|
match, _, _ = match_step(match, action)
|
||||||
total = np.asarray(match_score(match))
|
after = np.asarray(match_score(match))
|
||||||
# The total never resets when a round rolls over; it only accumulates.
|
paid += ((after[0] - after[1]) - (before[0] - before[1])) / scale
|
||||||
assert total.shape == (2,)
|
|
||||||
prev = total
|
final = np.asarray(match_score(match))
|
||||||
assert prev.shape == (2,)
|
total_lead = (final[0] - final[1]) / scale
|
||||||
|
# Every point banked along the way, and nothing else.
|
||||||
|
assert paid == pytest.approx(total_lead, abs=1e-4)
|
||||||
|
|
||||||
|
|
||||||
# --- both seats ------------------------------------------------------------
|
# --- both seats ------------------------------------------------------------
|
||||||
|
|||||||
Reference in New Issue
Block a user