File size: 15,412 Bytes
78738de
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
"""PositionPlayEnv — curriculum stage 2a cho PoolCoach.

Task: cue + bi 1 + bi 2 (vị trí random, cách nhau >= 4R), vẫn 1 cú đánh.
Pot bi 1 sao cho bi cái dừng ở vị trí có CÚ THUẬN đánh tiếp bi 2.

Thiết kế chốt 13/07/2026 (Documents/PoolCoach_Stage2_PositionPlay_Design.md):
GIỮ NGUYÊN toàn bộ chuỗi reward stage 1 đã phá collapse (pot +1.0, scratch -0.3,
gated contact bonus, dense aim ghost-ball, shaping khoảng cách) — chỉ THÊM đúng
1 term position, GATED theo (pot && !scratch):

    reward += POS_COEF * Q(cue_final, b2_final)      # POS_COEF = 0.5

Bài học gated contact bonus (11/07): không gate là agent farm position mà bỏ pot.
Pot suông +1.0, pot + shape đẹp +1.5 → thứ tự ưu tiên không đổi.

Position quality Q ∈ [0, 1]:
    với mỗi lỗ p:  u = unit(b2 − c);  v = unit(p − b2);  cos_cut = dot(u, v)
    Q_angle = max(0, max_p cos_cut)     # lỗ thuận nhất; không lỗ nào dương → 0
    Q = Q_angle * 1 / (1 + d(c, b2) / POS_DIST_REF)
Chọn lỗ theo cos_cut lớn nhất — cùng logic _ghost_aim_point (KHÔNG lỗ gần nhất).
dist_factor phạt nhẹ cú xa (d=0.5m → 0.67; d=1m → 0.5; d=2m → 0.33).
2 bi → không có snooker, bỏ qua check obstruction.

Obs: Box[0,1]^6 = [cue_x, cue_y, b1_x, b1_y, b2_x, b2_y] / (w, l), clip [0,1].
Action: giữ nguyên Box[-1,1]^4 (phi / V0 / side / vert). Stage 2 mới là lúc
spin có việc → log |side|, |vert| vào info (figure spin-usage cho luận văn).

Edge cases (design doc §5):
    - Bi 1 không vào       → gate đóng, không position term (như stage 1)
    - Scratch (kể cả pot)  → scratch_penalty, không position term
    - Bi 2 bị đụng văng    → hợp lệ, Q tính trên vị trí bi 2 SAU simulate
    - Bi 2 vô tình rớt lỗ  → Q = 1.0 (combo may mắn), log b2_potted theo dõi
    - simulation_failed    → như stage 1, info đủ key (pos_q=0, b2_potted=0)
"""

from __future__ import annotations

import numpy as np

import gymnasium as gym
from gymnasium import spaces

import pooltool as pt
import pooltool.constants as ptc

BALL_R = 0.028575  # bán kính bi (m), mặc định pooltool

# --- Reward config: chuỗi stage 1 giữ NGUYÊN (không đụng vào) ---
POT_REWARD = 1.0        # bi 1 vào lỗ
SCRATCH_PENALTY = -0.3  # bi cái vào lỗ — kể cả khi bi 1 cũng vào
CONTACT_BONUS = 0.2     # gated theo năng lượng truyền (chống tap farming)
CONTACT_DISP_REF = 0.5  # (m) quãng đường bi 1 để ăn TRỌN contact bonus
AIM_COEF = 0.1          # dense aim reward về ghost-ball của BI 1 (không đổi)

# --- Term MỚI duy nhất của stage 2a ---
POS_COEF = 0.5          # hệ số position term, GATED theo (pot && !scratch)
# Smoke test hỏi đúng 2 câu: pot% giữ ~15-20%? Q có tăng?
# Pot sập → giảm 0.5 → 0.25 (position term lấn multi-objective).
POS_DIST_REF = 1.0      # (m) tham chiếu dist_factor: d=1m → factor 0.5


def table_bounds(w, l, pockets):
    """Hộp bao của (mặt bàn ∪ mọi đường tròn lỗ) — trả (lo_x, hi_x, lo_y, hi_y).

    ``pockets``: iterable object có ``.center`` (≥2 chiều) và ``.radius``, tức
    đúng shape của ``pt.Table.pockets.values()``.

    Đây là ENVELOPE VẬT LÝ của mọi vị trí nghỉ hợp lệ: bi CHƯA bị ăn không thể
    nằm ngoài hộp này, vì muốn ra xa hơn nó phải đi xuyên qua một đường tròn
    lỗ, mà đi qua đường tròn lỗ là bị ăn. Tâm lỗ nằm NGOÀI hộp bàn (lỗ góc lệch
    ``depth/√2`` theo đường chéo, lỗ giữa lệch ``depth`` theo phương ngang) nên
    hộp này rộng hơn ``[0, w] × [0, l]`` thật sự, không phải chỉ trên lý thuyết.

    Tách thành hàm THUẦN (không đụng pooltool, không đụng self) có chủ đích:
    envelope phải KIỂM ĐƯỢC là dẫn ra từ hình học chứ không phải hằng số gõ
    tay — mà `tests/conftest.py` cố ý nhét stub pooltool rỗng nên test không
    dựng nổi bàn thật. Hàm thuần thì test nạp bàn giả vào được.
    """
    return (
        min(0.0, min(float(p.center[0]) - float(p.radius) for p in pockets)),
        max(float(w), max(float(p.center[0]) + float(p.radius)
                          for p in pockets)),
        min(0.0, min(float(p.center[1]) - float(p.radius) for p in pockets)),
        max(float(l), max(float(p.center[1]) + float(p.radius)
                          for p in pockets)),
    )


class PositionPlayEnv(gym.Env):
    """Một cú đánh: pot bi 1 + điều bi cái về vị trí thuận cho bi 2."""

    metadata = {"render_modes": []}

    def __init__(
        self,
        seed: int | None = None,
        scratch_penalty: float = SCRATCH_PENALTY,
        pos_coef: float = POS_COEF,
        aim_mode: str = "best_cut",
    ):
        """pos_coef: override cho ablation POS_COEF ∈ {0, 0.5, 1.0};
        pos_coef=0 chính là stage-1-trên-env-mới (position-blind baseline).

        aim_mode:
            "best_cut" (mặc định, = stage 1): aim reward về ghost-ball của
                MỘT lỗ có cut angle thuận nhất — phi bị ghim vào 1 đường pot.
            "any": aim reward = max cos trên ghost-ball của MỌI lỗ khả thi
                (dot(cue→bi, bi→lỗ) > 0) — agent tự do chọn lỗ, mở khoá phi
                cho position play (thí nghiệm 14/07: aim best_cut nghi là
                nút thắt khiến Q phẳng — phi khoá thì chỉ còn V0/spin tự do).
        """
        super().__init__()
        assert aim_mode in ("best_cut", "any"), aim_mode
        self.scratch_penalty = float(scratch_penalty)
        self.pos_coef = float(pos_coef)
        self.aim_mode = aim_mode
        self.table = pt.Table.default()
        self.w, self.l = self.table.w, self.table.l
        self._diag = float(np.hypot(self.w, self.l))
        self._pockets = [
            np.asarray(p.center[:2], dtype=np.float64)
            for p in self.table.pockets.values()
        ]
        # Envelope vị trí nghỉ hợp lệ, tính MỘT LẦN lúc dựng bàn — `validate_full`
        # chỉ đọc thuộc tính này, không tính lại mỗi cú. Xem `table_bounds`.
        self._bounds = table_bounds(self.w, self.l,
                                    self.table.pockets.values())

        self.observation_space = spaces.Box(0.0, 1.0, shape=(6,), dtype=np.float32)
        self.action_space = spaces.Box(-1.0, 1.0, shape=(4,), dtype=np.float32)
        self._rng = np.random.default_rng(seed)
        self._system: pt.System | None = None

    # ------------------------------------------------------------------ utils

    def _random_xy(self) -> np.ndarray:
        margin = 4 * BALL_R
        return np.array(
            [
                self._rng.uniform(margin, self.w - margin),
                self._rng.uniform(margin, self.l - margin),
            ]
        )

    def _min_pocket_dist(self, xy: np.ndarray) -> float:
        return min(float(np.linalg.norm(xy - p)) for p in self._pockets)

    def _ghost_aim_point(self, cue_xy: np.ndarray, tgt_xy: np.ndarray) -> np.ndarray:
        """Điểm ghost-ball cho BI 1 — copy nguyên từ stage 1 (không đổi)."""
        to_tgt = tgt_xy - cue_xy
        to_tgt = to_tgt / np.linalg.norm(to_tgt)
        best_dir, best_score = None, -np.inf
        for p in self._pockets:
            to_pocket = p - tgt_xy
            norm = float(np.linalg.norm(to_pocket))
            if norm < 1e-9:
                continue
            to_pocket = to_pocket / norm
            score = float(np.dot(to_tgt, to_pocket))
            if score > best_score:
                best_dir, best_score = to_pocket, score
        return tgt_xy - 2.0 * BALL_R * best_dir

    def _ghost_dirs_any(self, cue_xy: np.ndarray, tgt_xy: np.ndarray) -> list:
        """Hướng cue→ghost cho MỌI lỗ khả thi (aim_mode="any").

        Lỗ khả thi: dot(cue→bi, bi→lỗ) > 0 — ghost nằm phía bi cái với tới
        được. Lỗ ngược hướng bị loại (ghost sau lưng bi, cú bất khả thi;
        thưởng aim vào đó là thưởng cú trượt). Không lỗ nào khả thi (hiếm)
        → fallback về best_cut.
        """
        to_tgt = tgt_xy - cue_xy
        to_tgt = to_tgt / np.linalg.norm(to_tgt)
        dirs = []
        for p in self._pockets:
            v = p - tgt_xy
            norm = float(np.linalg.norm(v))
            if norm < 1e-9:
                continue
            v = v / norm
            if float(np.dot(to_tgt, v)) <= 0.0:
                continue
            g = (tgt_xy - 2.0 * BALL_R * v) - cue_xy
            g_norm = float(np.linalg.norm(g))
            if g_norm < 1e-9:
                continue
            dirs.append(g / g_norm)
        if not dirs:
            g = self._ghost_aim_point(cue_xy, tgt_xy) - cue_xy
            dirs.append(g / np.linalg.norm(g))
        return dirs

    def _position_q(self, cue_xy: np.ndarray, b2_xy: np.ndarray) -> float:
        """Q ∈ [0,1]: chất lượng vị trí bi cái cho cú tiếp theo vào bi 2.

        Q_angle = max(0, max_p cos_cut) — lỗ thuận nhất theo cut angle,
        cùng triết lý chọn lỗ với _ghost_aim_point.
        dist_factor = 1/(1 + d/POS_DIST_REF) — phạt nhẹ cú xa.
        """
        d = float(np.linalg.norm(b2_xy - cue_xy))
        if d < 1e-9:  # trùng vị trí (không xảy ra thực tế)
            return 0.0
        u = (b2_xy - cue_xy) / d  # hướng cú đánh bi 2
        best_cos = -np.inf
        for p in self._pockets:
            to_pocket = p - b2_xy
            norm = float(np.linalg.norm(to_pocket))
            if norm < 1e-9:
                continue
            best_cos = max(best_cos, float(np.dot(u, to_pocket / norm)))
        q_angle = max(0.0, best_cos)
        dist_factor = 1.0 / (1.0 + d / POS_DIST_REF)
        return q_angle * dist_factor

    def _ball_xy(self, ball_id: str) -> np.ndarray:
        return np.asarray(
            self._system.balls[ball_id].state.rvw[0][:2], dtype=np.float64
        )

    def _obs(self) -> np.ndarray:
        cue, b1, b2 = self._ball_xy("cue"), self._ball_xy("1"), self._ball_xy("2")
        obs = np.array(
            [
                cue[0] / self.w, cue[1] / self.l,
                b1[0] / self.w, b1[1] / self.l,
                b2[0] / self.w, b2[1] / self.l,
            ],
            dtype=np.float32,
        )
        # bi đã vào lỗ nằm ở tâm lỗ (hơi ngoài mép bàn) -> clip về [0, 1]
        return np.clip(obs, 0.0, 1.0)

    @staticmethod
    def _pocketed(system: pt.System, ball_id: str) -> bool:
        return system.balls[ball_id].state.s == ptc.pocketed

    @staticmethod
    def _cue_hit_target(system: pt.System) -> bool:
        """True nếu có va chạm bi-bi giữa 'cue' và '1' trong lượt mô phỏng."""
        for ev in pt.events.filter_type(system.events, pt.EventType.BALL_BALL):
            if "cue" in ev.ids and "1" in ev.ids:
                return True
        return False

    # ------------------------------------------------------------------ gym API

    def reset(self, *, seed: int | None = None, options=None):
        super().reset(seed=seed)
        if seed is not None:
            self._rng = np.random.default_rng(seed)

        # 3 bi, đôi một cách nhau > 4R
        placed = [self._random_xy()]
        while len(placed) < 3:
            xy = self._random_xy()
            if all(np.linalg.norm(xy - q) > 4 * BALL_R for q in placed):
                placed.append(xy)
        cue_xy, b1_xy, b2_xy = placed

        balls = {
            "cue": pt.Ball.create("cue", xy=tuple(cue_xy)),
            "1": pt.Ball.create("1", xy=tuple(b1_xy)),
            "2": pt.Ball.create("2", xy=tuple(b2_xy)),
        }
        self._system = pt.System(
            table=self.table, balls=balls, cue=pt.Cue(cue_ball_id="cue")
        )
        return self._obs(), {}

    def step(self, action):
        action = np.clip(np.asarray(action, dtype=np.float64), -1.0, 1.0)
        phi = float((action[0] + 1.0) / 2.0 * 360.0) % 360.0
        v0 = float(0.5 + (action[1] + 1.0) / 2.0 * 3.5)
        side = float(action[2] * 0.4)
        vert = float(action[3] * 0.4)
        abs_side, abs_vert = abs(side) / 0.4, abs(vert) / 0.4  # chuẩn hoá [0,1]

        # --- dense aim reward: về ghost-ball BI 1, tính TRƯỚC simulate ---
        cue_xy, tgt_xy = self._ball_xy("cue"), self._ball_xy("1")
        phi_rad = np.radians(phi)
        aim_dir = np.array([np.cos(phi_rad), np.sin(phi_rad)])
        if self.aim_mode == "any":
            # max trên mọi lỗ khả thi — agent tự chọn lỗ, phi không bị ghim
            aim_cos = max(
                float(np.dot(aim_dir, d))
                for d in self._ghost_dirs_any(cue_xy, tgt_xy)
            )
        else:  # "best_cut" — nguyên bản stage 1
            ghost = self._ghost_aim_point(cue_xy, tgt_xy)
            to_ghost = ghost - cue_xy
            to_ghost /= np.linalg.norm(to_ghost)  # |cue-ghost| >= 2R > 0
            aim_cos = float(np.dot(aim_dir, to_ghost))
        aim_r = AIM_COEF * aim_cos

        d_before = self._min_pocket_dist(tgt_xy)

        self._system.cue.set_state(V0=v0, phi=phi, a=side, b=vert)
        try:
            pt.simulate(self._system, inplace=True)
        except Exception:  # hiếm: edge case vật lý -> coi như foul
            return self._obs(), self.scratch_penalty + aim_r, True, False, {
                "potted": 0, "scratch": 0, "contact": 0, "aim_cos": aim_cos,
                "tgt_disp": 0.0, "pos_q": 0.0, "b2_potted": 0,
                "abs_side": abs_side, "abs_vert": abs_vert,
                "error": "simulation_failed",
            }

        scratch = self._pocketed(self._system, "cue")
        potted = self._pocketed(self._system, "1")
        b2_potted = self._pocketed(self._system, "2")
        contact = self._cue_hit_target(self._system)
        tgt_disp = float(np.linalg.norm(self._ball_xy("1") - tgt_xy))

        # --- position quality: chỉ tính khi gate mở (pot && !scratch) ---
        pos_q = 0.0
        if potted and not scratch:
            if b2_potted:
                pos_q = 1.0  # combo may mắn — theo dõi tần suất qua b2_potted
            else:
                pos_q = self._position_q(self._ball_xy("cue"), self._ball_xy("2"))

        if scratch:
            reward = self.scratch_penalty
        elif potted:
            reward = POT_REWARD + self.pos_coef * pos_q
        else:
            d_after = self._min_pocket_dist(self._ball_xy("1"))
            reward = (d_before - d_after) / self._diag
            if contact:
                reward += CONTACT_BONUS * min(1.0, tgt_disp / CONTACT_DISP_REF)
        reward += aim_r  # áp lên MỌI cú — no-op không né được

        # ép int/float để VecMonitor(info_keywords=...) ghi vào monitor.csv
        info = {
            "potted": int(potted), "scratch": int(scratch),
            "contact": int(contact), "aim_cos": aim_cos, "tgt_disp": tgt_disp,
            "pos_q": float(pos_q), "b2_potted": int(b2_potted),
            "abs_side": abs_side, "abs_vert": abs_vert,
        }
        return self._obs(), float(reward), True, False, info