Spaces:
Sleeping
Sleeping
| # -*- coding: utf-8 -*- | |
| """Unit test ShotNet (BG27 bước 1) — loader shape/mask trên shard THẬT, | |
| forward pass, angular loss quanh 0°/360°, mask loss (a,b) trên cú | |
| non-identifiable. | |
| Chạy được trong venv app (torch CPU có sẵn — cố ý của workspace); không cần | |
| GPU, không cần pooltool (loader đọc npz thẳng). Test shard thật skip nếu | |
| máy không có ``datasets\\bb9_synth`` (dataset ngoài git). | |
| """ | |
| from __future__ import annotations | |
| from pathlib import Path | |
| import numpy as np | |
| import pytest | |
| import torch | |
| from poolcoach_cv import shotnet as sn | |
| DATA_DIR = Path(r"D:\Khoa luan\datasets\bb9_synth") | |
| HELDOUT = DATA_DIR / "heldout_0000000.npz" | |
| TINY = sn.ShotNetConfig(d_model=32, n_layers=2, n_heads=2, d_ff=64, | |
| dropout=0.0) | |
| # ------------------------------------------------------------ loader thật | |
| def real_ds(): | |
| if not HELDOUT.exists(): | |
| pytest.skip("khong co dataset bb9_synth tren may nay") | |
| return sn.ShardDataset([HELDOUT], limit_shots=50) | |
| def test_loader_shapes_on_real_shard(real_ds): | |
| assert len(real_ds) == 50 | |
| for i in (0, 7, 49): | |
| shot = real_ds.raw_shot(i) | |
| item = real_ds[i] | |
| F_n = shot["n_frames"] | |
| assert item["feats"].shape == (F_n, sn.FEAT_DIM) | |
| assert item["t"].shape == (F_n,) | |
| assert real_ds.lengths[i] == F_n | |
| # t dựng lại từ fps (PTS CFR — format dataset) | |
| np.testing.assert_allclose( | |
| item["t"], np.arange(F_n) / shot["fps"], rtol=0, atol=1e-6) | |
| # cột vis khớp covered từng bi (mask đúng — yêu cầu BRIEF 1.4) | |
| vis_cols = item["feats"][:, 2::3][:, :sn.N_SLOTS] | |
| assert vis_cols.sum() == shot["covered"].sum() | |
| # vị trí chuẩn hoá đẳng hướng nằm quanh bàn (nhiễu cho phép lệch nhẹ) | |
| assert np.abs(item["feats"][:, 0:-1]).max() < 1.5 | |
| # img_diff đã clip [0, 2] (sentinel −1 của frame đầu về 0) | |
| assert item["feats"][0, -1] == 0.0 | |
| assert item["feats"][:, -1].min() >= 0.0 | |
| assert item["feats"][:, -1].max() <= sn.IMG_DIFF_CLIP | |
| def test_loader_labels_on_real_shard(real_ds): | |
| z = np.load(HELDOUT) | |
| for i in (0, 3): | |
| item = real_ds[i] | |
| for k in ("label_v0", "label_phi", "label_a", "label_b", | |
| "identifiable", "v0_ball", "fps"): | |
| assert item[k] == pytest.approx(float(z[k][i]), abs=1e-6) | |
| def test_loader_khop_iter_shots(real_ds): | |
| """Đường train (ShardDataset.raw_shot) phải trả ĐÚNG dữ liệu của đường | |
| eval (gen_synth_shots.iter_shots) — hai loader một nguồn, lệch là net | |
| train một đằng chấm một nẻo.""" | |
| import sys | |
| sys.path.insert(0, str(Path(__file__).resolve().parents[1] | |
| / "scripts" / "broadcast")) | |
| from gen_synth_shots import iter_shots | |
| for i, ref in enumerate(iter_shots(HELDOUT)): | |
| if i >= 5: | |
| break | |
| mine = real_ds.raw_shot(i) | |
| np.testing.assert_array_equal(mine["xy"], ref["xy"]) | |
| np.testing.assert_array_equal(mine["covered"], ref["covered"]) | |
| np.testing.assert_array_equal(mine["img_diff"], ref["img_diff"]) | |
| np.testing.assert_array_equal(mine["ball_ids"], ref["ball_ids"]) | |
| np.testing.assert_allclose(mine["t"], ref["t"], atol=1e-6) | |
| assert mine["label_v0"] == pytest.approx(ref["label_v0"]) | |
| assert mine["identifiable"] == ref["identifiable"] | |
| # ------------------------------------------------- featurize (cú nhân tạo) | |
| def _toy_shot(): | |
| """2 bi (cue + bi 5), 3 frame, số chọn tay để soi từng ô feature.""" | |
| w, l = sn.TABLE_W_M, sn.TABLE_L_M | |
| xy = np.array([[[w / 2, l / 2], [w / 2 + 0.1, l / 2]], | |
| [[w / 2, l / 2 + 0.2], [w / 2 + 0.1, l / 2]], | |
| [[0.0, 0.0], [w / 2 + 0.1, l / 2]]], dtype=np.float32) | |
| covered = np.array([[True, True], [True, False], [False, True]]) | |
| return {"xy": xy, "covered": covered, | |
| "img_diff": np.array([-1.0, 0.5, 3.0], dtype=np.float32), | |
| "ball_ids": np.array([0, 5], dtype=np.uint8), | |
| "t": np.array([0.0, 0.1, 0.2], dtype=np.float32)} | |
| def test_featurize_slot_mapping_and_mask(): | |
| feats, t = sn.featurize_shot(_toy_shot()) | |
| assert feats.shape == (3, sn.FEAT_DIM) | |
| s = 2.0 / sn.TABLE_L_M | |
| # frame 0: cue giữa bàn → slot 0 = (0, 0, vis 1); bi 5 lệch x +0.1 | |
| np.testing.assert_allclose(feats[0, 0:3], [0.0, 0.0, 1.0], atol=1e-6) | |
| np.testing.assert_allclose(feats[0, 15:18], [0.1 * s, 0.0, 1.0], | |
| atol=1e-6) | |
| # frame 1: cue nhích y +0.2; bi 5 KHÔNG thấy → (0, 0, 0) | |
| np.testing.assert_allclose(feats[1, 0:3], [0.0, 0.2 * s, 1.0], atol=1e-6) | |
| np.testing.assert_allclose(feats[1, 15:18], [0.0, 0.0, 0.0], atol=1e-6) | |
| # frame 2: cue KHÔNG thấy → 0 dù toạ độ thô là góc bàn | |
| np.testing.assert_allclose(feats[2, 0:3], [0.0, 0.0, 0.0], atol=1e-6) | |
| # slot không có bi trên bàn (vd bi 1) = 0 tuyệt đối | |
| assert np.all(feats[:, 3:15] == 0.0) | |
| # img_diff: sentinel −1 → 0; 0.5 giữ; 3.0 clip về 2.0 | |
| np.testing.assert_allclose(feats[:, -1], [0.0, 0.5, 2.0], atol=1e-6) | |
| def test_featurize_deltas_mode(): | |
| """Chế độ deltas (config c2): dx/dy = hiệu vị trí frame trước, chỉ khi | |
| CẢ HAI frame thấy bi; frame đầu = 0; img_diff vẫn là cột cuối.""" | |
| feats, _t = sn.featurize_shot(_toy_shot(), deltas=True) | |
| assert feats.shape == (3, sn.FEAT_DIM_DELTAS) | |
| s = 2.0 / sn.TABLE_L_M | |
| d0 = 3 * sn.N_SLOTS # khối delta bắt đầu sau khối (x,y,vis) | |
| # frame 0: mọi delta = 0 | |
| assert np.all(feats[0, d0:-1] == 0.0) | |
| # frame 1: cue thấy ở cả 0 và 1, nhích y +0.2 → (0, 0.2·s) | |
| np.testing.assert_allclose(feats[1, d0:d0 + 2], [0.0, 0.2 * s], | |
| atol=1e-6) | |
| # bi 5 (slot 5) không thấy ở frame 1 → delta 0 | |
| np.testing.assert_allclose(feats[1, d0 + 10:d0 + 12], [0.0, 0.0], | |
| atol=1e-6) | |
| # frame 2: cue mất → delta 0; bi 5 thấy lại nhưng frame TRƯỚC không thấy | |
| # → delta vẫn 0 (không nhảy vọt qua gap) | |
| assert np.all(feats[2, d0:-1] == 0.0) | |
| # khối (x, y, vis) y hệt chế độ thường; img_diff cột cuối | |
| base, _ = sn.featurize_shot(_toy_shot()) | |
| np.testing.assert_array_equal(feats[:, :3 * sn.N_SLOTS], | |
| base[:, :3 * sn.N_SLOTS]) | |
| np.testing.assert_array_equal(feats[:, -1], base[:, -1]) | |
| def test_ab_scale_loss_and_predict_roundtrip(): | |
| """ab_scale (c2): loss chấm trong không gian (a,b)/scale — head dự đúng | |
| target scaled thì loss ab = 0; predict trả về ĐƠN VỊ GỐC.""" | |
| cfg = sn.ShotNetConfig(d_model=32, n_layers=1, n_heads=2, d_ff=64, | |
| dropout=0.0, ab_scale=0.4) | |
| b = _rand_batch() | |
| out = {"v0_z": torch.log(b["label_v0"]), | |
| "phi_vec": torch.stack( | |
| [torch.cos(torch.deg2rad(b["label_phi"])), | |
| torch.sin(torch.deg2rad(b["label_phi"]))], dim=-1), | |
| "ab": torch.stack([b["label_a"], b["label_b"]], dim=-1) / 0.4, | |
| "ident_logit": torch.where(b["identifiable"] > 0.5, 20.0, -20.0)} | |
| loss = sn.shotnet_loss(out, b, cfg) | |
| assert loss["ab"].item() == pytest.approx(0.0, abs=1e-8) | |
| assert loss["v0"].item() == pytest.approx(0.0, abs=1e-8) | |
| # predict nhân ngược ab_scale → đơn vị gốc | |
| torch.manual_seed(0) | |
| model = sn.ShotNet(cfg) | |
| with torch.no_grad(): | |
| raw = model(b["x"][:, :, :cfg.feat_dim], b["t"], b["mask"]) | |
| p = model.predict(b["x"][:, :, :cfg.feat_dim], b["t"], b["mask"]) | |
| np.testing.assert_allclose(p["a"], raw["ab"][:, 0].numpy() * 0.4, | |
| atol=1e-6) | |
| def test_config_deltas_feat_dim_tu_nang(): | |
| assert sn.ShotNetConfig(use_deltas=True).feat_dim == sn.FEAT_DIM_DELTAS | |
| assert sn.ShotNetConfig().feat_dim == sn.FEAT_DIM | |
| # ----------------------------------------------------- split + batch + pad | |
| def test_val_split_deterministic_disjoint(): | |
| tr1, va1 = sn.val_split_indices(1000, 0.02, seed=123) | |
| tr2, va2 = sn.val_split_indices(1000, 0.02, seed=123) | |
| np.testing.assert_array_equal(va1, va2) | |
| np.testing.assert_array_equal(tr1, tr2) | |
| assert len(va1) == 20 and len(tr1) == 980 | |
| assert np.intersect1d(tr1, va1).size == 0 | |
| assert not np.array_equal(sn.val_split_indices(1000, 0.02, 7)[1], va1) | |
| def test_bucket_batcher_covers_all_within_budget(): | |
| rng = np.random.default_rng(0) | |
| lengths = rng.integers(10, 1200, size=500) | |
| idx = np.arange(500) | |
| bb = sn.BucketBatcher(lengths, idx, token_budget=8_000, max_batch=64, | |
| seed=1) | |
| batches = bb.epoch_batches(epoch=0) | |
| got = np.sort(np.concatenate(batches)) | |
| np.testing.assert_array_equal(got, idx) | |
| for b in batches: | |
| assert len(b) == 1 or len(b) * lengths[b].max() <= 8_000 | |
| # tái lập theo (seed, epoch) | |
| again = bb.epoch_batches(epoch=0) | |
| assert all(np.array_equal(x, y) for x, y in zip(batches, again)) | |
| def test_collate_padding_and_mask(): | |
| def item(n): | |
| return {"feats": np.ones((n, sn.FEAT_DIM), np.float32), | |
| "t": np.arange(n, dtype=np.float32) / 30.0, | |
| **{k: 1.0 for k in sn.LABEL_KEYS}} | |
| out = sn.collate([item(3), item(5)]) | |
| assert out["x"].shape == (2, 5, sn.FEAT_DIM) | |
| assert out["mask"].tolist() == [[True] * 3 + [False] * 2, [True] * 5] | |
| assert torch.all(out["x"][0, 3:] == 0) | |
| assert out["label_v0"].shape == (2,) | |
| # -------------------------------------------------------- model + loss | |
| def _rand_batch(B=3, T=7, seed=0): | |
| g = torch.Generator().manual_seed(seed) | |
| return {"x": torch.randn(B, T, sn.FEAT_DIM, generator=g), | |
| "t": torch.arange(T).float().repeat(B, 1) / 30.0, | |
| "mask": torch.tensor([[True] * T, [True] * (T - 2) + [False] * 2, | |
| [True] * T]), | |
| "label_v0": torch.tensor([1.0, 3.0, 6.0]), | |
| "label_phi": torch.tensor([10.0, 200.0, 359.0]), | |
| "label_a": torch.tensor([0.2, -0.3, 0.0]), | |
| "label_b": torch.tensor([-0.2, 0.1, 0.3]), | |
| "identifiable": torch.tensor([1.0, 0.0, 1.0]), | |
| "v0_ball": torch.tensor([1.4, 4.0, 8.0]), | |
| "phi_ball": torch.tensor([11.0, 199.0, 358.0]), | |
| "fps": torch.tensor([30.0, 60.0, 25.0]), | |
| "upconvert": torch.zeros(3), "scratch": torch.zeros(3)} | |
| def test_forward_pass_shapes_finite(): | |
| torch.manual_seed(0) | |
| model = sn.ShotNet(TINY) | |
| b = _rand_batch() | |
| out = model(b["x"], b["t"], b["mask"]) | |
| assert out["v0_z"].shape == (3,) | |
| assert out["phi_vec"].shape == (3, 2) | |
| assert out["ab"].shape == (3, 2) | |
| assert out["ident_logit"].shape == (3,) | |
| for v in out.values(): | |
| assert torch.isfinite(v).all() | |
| loss = sn.shotnet_loss(out, b, TINY) | |
| assert torch.isfinite(loss["total"]) | |
| loss["total"].backward() # gradient chảy về input proj | |
| assert model.inp.weight.grad is not None | |
| def test_predict_ranges(): | |
| torch.manual_seed(0) | |
| model = sn.ShotNet(TINY) | |
| b = _rand_batch() | |
| p = model.predict(b["x"], b["t"], b["mask"]) | |
| assert np.all(p["v0"] > 0) # exp() — thước gậy m/s | |
| assert np.all((p["phi_deg"] >= 0) & (p["phi_deg"] < 360)) | |
| assert np.all((p["p_ident"] >= 0) & (p["p_ident"] <= 1)) | |
| def test_phi_loss_wraps_around_zero(): | |
| # pred 1° phải GẦN target 359° (Δ=2°), pred 181° phải XA (Δ=178°) | |
| target = torch.tensor([359.0]) | |
| near = torch.tensor([[np.cos(np.deg2rad(1.0)), np.sin(np.deg2rad(1.0))]], | |
| dtype=torch.float32) | |
| far = torch.tensor([[np.cos(np.deg2rad(181.0)), | |
| np.sin(np.deg2rad(181.0))]], dtype=torch.float32) | |
| l_near = sn._phi_vec_loss(near, target).item() | |
| l_far = sn._phi_vec_loss(far, target).item() | |
| assert l_near < 0.005 and l_far > 1.0 | |
| # metric vòng tròn cùng công thức harness | |
| assert sn.circ_diff_deg(359.0, 1.0) == pytest.approx(2.0) | |
| assert sn.circ_diff_deg(0.0, 360.0) == pytest.approx(0.0) | |
| def test_phi_ang_huber_wrap_va_chan_duoi(): | |
| """ang_huber (c3): vẫn liên tục quanh 0°/360°; gradient theo góc bị | |
| CHẶN TRẦN ngoài delta (cú 90° không được kéo mạnh hơn cú 9° — vec-MSE | |
| thì có, ∝ sin Δ).""" | |
| def vec(deg): | |
| return torch.tensor( | |
| [[np.cos(np.deg2rad(deg)), np.sin(np.deg2rad(deg))]], | |
| dtype=torch.float32, requires_grad=True) | |
| target = torch.tensor([359.0]) | |
| near, far = vec(1.0), vec(181.0) | |
| l_near = sn._phi_ang_huber(near, target, 5.0).sum() | |
| l_far = sn._phi_ang_huber(far, target, 5.0).sum() | |
| assert l_near.item() < 0.001 and l_far.item() > 0.1 | |
| # gradient theo góc xấp xỉ bằng nhau ở 30° và 120° (đều vùng linear) | |
| g = {} | |
| for name, deg in [("mid", 30.0), ("tail", 120.0)]: | |
| v = vec(deg) | |
| sn._phi_ang_huber(v, torch.tensor([0.0]), 5.0).sum().backward() | |
| g[name] = float(torch.linalg.vector_norm(v.grad)) | |
| assert g["mid"] == pytest.approx(g["tail"], rel=0.05) | |
| # config nối đúng loss: ang_huber phạt cú 90° nhẹ hơn vec_mse | |
| cfg_ang = sn.ShotNetConfig(phi_loss="ang_huber", w_phi=1.0) | |
| cfg_vec = sn.ShotNetConfig(phi_loss="vec_mse", w_phi=1.0) | |
| b = _rand_batch() | |
| out = {"v0_z": torch.log(b["label_v0"]), | |
| "phi_vec": vec(90.0).detach().repeat(3, 1), | |
| "ab": torch.zeros(3, 2), | |
| "ident_logit": torch.zeros(3)} | |
| b90 = dict(b) | |
| b90["label_phi"] = torch.tensor([0.0, 0.0, 0.0]) | |
| assert sn.shotnet_loss(out, b90, cfg_ang)["phi"].item() < \ | |
| sn.shotnet_loss(out, b90, cfg_vec)["phi"].item() | |
| def test_ab_loss_masked_on_non_identifiable(): | |
| torch.manual_seed(1) | |
| model = sn.ShotNet(TINY) | |
| b = _rand_batch() | |
| out = model(b["x"], b["t"], b["mask"]) | |
| base = sn.shotnet_loss(out, b, TINY)["ab"].item() | |
| # đổi label (a,b) của cú NON-identifiable (index 1) → loss ab không đổi | |
| b2 = dict(b) | |
| b2["label_a"] = b["label_a"].clone() | |
| b2["label_a"][1] = 0.39 | |
| b2["label_b"] = b["label_b"].clone() | |
| b2["label_b"][1] = -0.39 | |
| assert sn.shotnet_loss(out, b2, TINY)["ab"].item() == pytest.approx(base) | |
| # đổi label của cú identifiable (index 0) → loss ab PHẢI đổi | |
| b3 = dict(b) | |
| b3["label_a"] = b["label_a"].clone() | |
| b3["label_a"][0] = -0.39 | |
| assert sn.shotnet_loss(out, b3, TINY)["ab"].item() != pytest.approx(base) | |
| # ------------------------------- slot-builder sim→real (BG28 bước 2.2) | |
| # Khe input chính của BG28: net cần slot bi mục tiêu BỀN theo thời gian như | |
| # slot sim; slot rác (id nhảy loạn) là input lệch phân phối → output rác | |
| # không báo trước. 3 ca BRIEF bắt buộc: đủ bi · mất frame · hai bi lướt gần. | |
| FPS = 30.0 | |
| def _frames(n): | |
| return np.arange(n, dtype=np.float32) / FPS | |
| def test_slot_builder_du_bi(): | |
| """Ca 1 — đủ bi: 3 bi tĩnh thấy mọi frame → đúng 3 slot, vis toàn 1, | |
| slot đứng yên tại chỗ (id không nhảy).""" | |
| pts = [(0.3, 0.5), (0.9, 1.2), (0.6, 2.0)] | |
| F = 10 | |
| dets = [list(pts) for _ in range(F)] | |
| xy, vis = sn.build_target_slots(_frames(F), dets) | |
| assert xy.shape == (F, 3, 2) and vis.shape == (F, 3) | |
| assert vis.all() | |
| for k, p in enumerate(pts): # thứ tự slot = thứ tự xuất hiện | |
| np.testing.assert_allclose(xy[:, k], np.tile(p, (F, 1)), atol=1e-6) | |
| def test_slot_builder_mat_frame(): | |
| """Ca 2 — mất frame: bi biến mất 3 frame giữa chừng → vis 0 đúng chỗ, | |
| quay lại vẫn NHẬP CÙNG slot cũ (không mọc slot mới).""" | |
| F = 12 | |
| dets = [] | |
| for f in range(F): | |
| row = [(0.3, 0.5)] | |
| if not (4 <= f <= 6): | |
| row.append((1.0, 1.5)) | |
| dets.append(row) | |
| xy, vis = sn.build_target_slots(_frames(F), dets) | |
| assert xy.shape[1] == 2 # vẫn đúng 2 slot — không mọc slot 3 | |
| np.testing.assert_array_equal( | |
| vis[:, 1], [f < 4 or f > 6 for f in range(F)]) | |
| assert vis[:, 0].all() | |
| # bi vào lỗ (mất hẳn từ frame 8): mask 0 tới hết, slot không bị tái dụng | |
| dets2 = [[(0.3, 0.5)] + ([(1.0, 1.5)] if f < 8 else []) for f in range(F)] | |
| _xy2, vis2 = sn.build_target_slots(_frames(F), dets2) | |
| np.testing.assert_array_equal(vis2[:, 1], [f < 8 for f in range(F)]) | |
| def test_slot_builder_hai_bi_luot_gan(): | |
| """Ca 3 — hai bi lướt gần (BRIEF bắt buộc): hai bi chạy ngược chiều, | |
| lúc sát nhất cách 0.04m (> ngưỡng dedup 0.03) → greedy NN không được | |
| tráo slot: mỗi slot giữ nguyên tuyến y của bi mình suốt chuỗi.""" | |
| F = 19 | |
| t = _frames(F) | |
| xa = np.linspace(0.2, 0.8, F) # bi A: trái → phải, y = 0.50 | |
| xb = np.linspace(0.8, 0.2, F) # bi B: phải → trái, y = 0.54 | |
| dets = [[(float(xa[f]), 0.50), (float(xb[f]), 0.54)] for f in range(F)] | |
| xy, vis = sn.build_target_slots(t, dets) | |
| assert xy.shape[1] == 2 and vis.all() | |
| np.testing.assert_allclose(xy[:, 0, 1], 0.50, atol=1e-6) # không tráo | |
| np.testing.assert_allclose(xy[:, 1, 1], 0.54, atol=1e-6) | |
| np.testing.assert_allclose(xy[:, 0, 0], xa, atol=1e-6) | |
| np.testing.assert_allclose(xy[:, 1, 0], xb, atol=1e-6) | |
| def test_slot_builder_dedup_double_detect(): | |
| """Double-detect (2 det < 1R cùng frame) không được đẻ slot ma.""" | |
| F = 5 | |
| dets = [[(0.5, 0.5), (0.51, 0.5)] for _ in range(F)] # cách 1cm < 0.03 | |
| xy, vis = sn.build_target_slots(_frames(F), dets) | |
| assert xy.shape[1] == 1 | |
| assert vis.all() | |
| def test_shot_from_track_dung_shape_loader(): | |
| """shot_from_track: rows/others của analyze_track → dict cú đúng shape | |
| featurize_shot cần; cue = slot 0, t trừ mốc frame đầu, img_diff đi | |
| nguyên; featurize với bàn THẬT (broadcast 1.27×2.54) cho toạ độ chuẩn | |
| hoá đẳng hướng đúng công thức.""" | |
| w, l = 1.27, 2.54 | |
| rows = [ | |
| {"frame_file": "00000", "t_s": 10.0, "covered": 1, | |
| "table_x_m": w / 2, "table_y_m": l / 2, "img_diff": -1.0}, | |
| {"frame_file": "00001", "t_s": 10.0 + 1 / FPS, "covered": 1, | |
| "table_x_m": w / 2, "table_y_m": l / 2 + 0.2, "img_diff": 0.5}, | |
| {"frame_file": "00002", "t_s": 10.0 + 2 / FPS, "covered": 0, | |
| "table_x_m": "", "table_y_m": "", "img_diff": 3.0}, | |
| ] | |
| others = [{"t_s": 10.0, "x_m": 0.3, "y_m": 0.5}, | |
| {"t_s": 10.0 + 2 / FPS, "x_m": 0.3, "y_m": 0.5}] | |
| shot = sn.shot_from_track(rows, others) | |
| assert shot["n_frames"] == 3 and shot["n_balls"] == 2 | |
| np.testing.assert_allclose(shot["t"], [0.0, 1 / FPS, 2 / FPS], atol=1e-6) | |
| np.testing.assert_array_equal(shot["ball_ids"], [0, 1]) | |
| np.testing.assert_array_equal(shot["covered"], | |
| [[True, True], [True, False], | |
| [False, True]]) | |
| feats, t = sn.featurize_shot(shot, w=w, l=l) | |
| assert feats.shape == (3, sn.FEAT_DIM) | |
| s = 2.0 / l | |
| # frame 0: cue giữa bàn → slot 0 = (0, 0, 1); bi slot 1 lệch tâm | |
| np.testing.assert_allclose(feats[0, 0:3], [0.0, 0.0, 1.0], atol=1e-6) | |
| np.testing.assert_allclose( | |
| feats[0, 3:6], [(0.3 - w / 2) * s, (0.5 - l / 2) * s, 1.0], | |
| atol=1e-6) | |
| # frame 1: cue nhích y +0.2 (đẳng hướng theo l); bi mất frame → 0 | |
| np.testing.assert_allclose(feats[1, 0:3], [0.0, 0.2 * s, 1.0], atol=1e-5) | |
| np.testing.assert_allclose(feats[1, 3:6], [0.0, 0.0, 0.0], atol=1e-6) | |
| # img_diff: sentinel −1 → 0, 0.5 giữ, 3.0 clip 2.0 (cột cuối) | |
| np.testing.assert_allclose(feats[:, -1], [0.0, 0.5, 2.0], atol=1e-6) | |
| # -------------------- augmentation "detections thưa" (BG29 bước 1.3) | |
| # Chẩn đoán BG28: net sập trên clip thật vì PHÂN PHỐI mật độ detection bi | |
| # mục tiêu (cú 11 = 1.01 det không-cue/frame, cú 12 = 3.09, synth 95–96% | |
| # visibility ≈ 5.3 det/frame). 4 ca BRIEF bắt buộc: seed tái lập · cue KHÔNG | |
| # rơi · keep-rate áp đúng phân phối · cú không augment bit-giống đường cũ. | |
| def _aug_shot(F_n=200, n_target=6): | |
| """Cú giả nhiều frame/nhiều bi — đủ mẫu để đo tỉ lệ sống sót.""" | |
| B = n_target + 1 | |
| return {"xy": np.zeros((F_n, B, 2), dtype=np.float32), | |
| "covered": np.ones((F_n, B), dtype=bool), | |
| "img_diff": np.zeros(F_n, dtype=np.float32), | |
| "ball_ids": np.arange(B, dtype=np.int64), | |
| "t": np.arange(F_n, dtype=np.float32) / 30.0, | |
| "t_first_bb": 0.5, "t_first_cush": float("nan")} | |
| def test_dropout_khong_cham_cue_va_khong_sua_tai_cho(): | |
| cfg = sn.SlotDropoutConfig(p_apply=1.0, keep_min=0.2, keep_max=0.2) | |
| shot = _aug_shot() | |
| orig = shot["covered"].copy() | |
| out, keep = sn.apply_slot_dropout(shot, cfg, | |
| np.random.default_rng([1, 2, 3])) | |
| assert keep == pytest.approx(0.2) | |
| assert out["covered"][:, 0].all() # cue KHÔNG rơi frame nào | |
| assert not out["covered"][:, 1:].all() # bi mục tiêu CÓ rơi | |
| # mảng gốc không bị ghi đè (raw_shot trả view vào shard nạp sẵn RAM) | |
| np.testing.assert_array_equal(shot["covered"], orig) | |
| def test_dropout_keep_rate_dung_phan_phoi(): | |
| """keep-rate cố định → tỉ lệ detection sống sót của bi mục tiêu bằng | |
| đúng keep-rate (sai số thống kê ~1/sqrt(200·6)).""" | |
| for keep in (0.15, 0.5, 0.9): | |
| cfg = sn.SlotDropoutConfig(p_apply=1.0, keep_min=keep, keep_max=keep) | |
| out, k = sn.apply_slot_dropout(_aug_shot(), cfg, | |
| np.random.default_rng([7, int(keep * 100)])) | |
| assert k == pytest.approx(keep) | |
| assert out["covered"][:, 1:].mean() == pytest.approx(keep, abs=0.04) | |
| def test_dropout_dai_keep_phu_vung_that(): | |
| """Dải keep-rate marginal phải PHỦ vùng đã đo trên clip thật: từ ~15% | |
| visibility (cú 11) tới 95–96% (synth gốc, nhánh không augment = 1.0).""" | |
| cfg = sn.SlotDropoutConfig() | |
| rng = np.random.default_rng(0) | |
| keeps = np.array([sn.apply_slot_dropout(_aug_shot(4, 3), cfg, rng)[1] | |
| for _ in range(3000)]) | |
| assert keeps.min() < 0.12 and keeps.min() >= cfg.keep_min | |
| assert (keeps == 1.0).mean() == pytest.approx(1 - cfg.p_apply, abs=0.03) | |
| aug = keeps[keeps < 1.0] | |
| assert aug.max() <= cfg.keep_max | |
| for lo, hi in ((0.10, 0.20), (0.40, 0.60), (0.85, 0.95)): | |
| assert ((aug >= lo) & (aug < hi)).sum() > 0 # phủ cả 3 vùng đo được | |
| def test_dropout_cu_khong_co_bi_muc_tieu_di_nguyen(): | |
| shot = _aug_shot(10, 0) # chỉ có cue | |
| out, keep = sn.apply_slot_dropout(shot, sn.SlotDropoutConfig(p_apply=1.0), | |
| np.random.default_rng(0)) | |
| assert keep == 1.0 and out is shot | |
| def test_loader_augment_mac_dinh_TAT_va_bit_giong_duong_cu(real_ds): | |
| """Ca quan trọng nhất: mặc định TẮT, và cú KHÔNG augment (p_apply=0) | |
| phải đi qua loader ra feats **bit giống** đường cũ — mọi đường eval | |
| (harness held-out, worker app, val mỗi epoch) dựa vào điều này.""" | |
| base = [real_ds[i]["feats"].copy() for i in (0, 5, 17)] | |
| ds = sn.ShardDataset([HELDOUT], limit_shots=50, | |
| augment=sn.SlotDropoutConfig(p_apply=1.0, | |
| keep_min=0.1, | |
| keep_max=0.1), | |
| reweight=sn.TargetReweight()) | |
| # chưa bật train mode → augment/reweight KHÔNG có hiệu lực | |
| for k, i in enumerate((0, 5, 17)): | |
| np.testing.assert_array_equal(ds[i]["feats"], base[k]) | |
| assert "sample_w" not in ds[i] | |
| # bật train mode với p_apply=0 → vẫn bit giống, nhưng sample_w xuất hiện | |
| ds0 = sn.ShardDataset([HELDOUT], limit_shots=50, | |
| augment=sn.SlotDropoutConfig(p_apply=0.0)) | |
| ds0.set_train_mode(True, epoch=3) | |
| for k, i in enumerate((0, 5, 17)): | |
| np.testing.assert_array_equal(ds0[i]["feats"], base[k]) | |
| assert ds0[i]["keep_rate"] == 1.0 | |
| def test_loader_augment_tai_lap_theo_seed_epoch(real_ds): | |
| """Cùng (seed, epoch, index) → mask y hệt; khác epoch → mask khác | |
| (augmentation thật sự on-the-fly, không phải một bản thưa cố định).""" | |
| def build(): | |
| return sn.ShardDataset([HELDOUT], limit_shots=50, | |
| augment=sn.SlotDropoutConfig(seed=99)) | |
| a, b = build(), build() | |
| a.set_train_mode(True, epoch=2) | |
| b.set_train_mode(True, epoch=2) | |
| for i in (0, 1, 2, 3, 4, 5, 6, 7): | |
| np.testing.assert_array_equal(a[i]["feats"], b[i]["feats"]) | |
| c = build() | |
| c.set_train_mode(True, epoch=3) | |
| khac = sum(not np.array_equal(a[i]["feats"], c[i]["feats"]) | |
| for i in range(20)) | |
| assert khac >= 10 # đa số cú đổi mask khi sang epoch khác | |
| # cột vis của cue (slot 0) không bao giờ bị augment chạm vào | |
| for i in range(20): | |
| np.testing.assert_array_equal(a[i]["feats"][:, 2], | |
| real_ds[i]["feats"][:, 2]) | |
| # augment CHỈ làm THƯA: mọi detection sống sót phải là detection cũ | |
| for i in range(20): | |
| vis_a = a[i]["feats"][:, 2::3][:, :sn.N_SLOTS] | |
| vis_0 = real_ds[i]["feats"][:, 2::3][:, :sn.N_SLOTS] | |
| assert np.all(vis_a <= vis_0) | |
| def test_reweight_lat_target_dung_nguong(real_ds): | |
| ds = sn.ShardDataset([HELDOUT], limit_shots=50, | |
| reweight=sn.TargetReweight(weight=3.0)) | |
| ds.set_train_mode(True) | |
| n_tgt = 0 | |
| for i in range(50): | |
| shot = ds.raw_shot(i) | |
| tf = sn.t_first_contact_s(shot) | |
| trong_lat = (not np.isnan(tf)) and tf >= sn.T_FIRST_TARGET_S | |
| assert ds[i]["sample_w"] == (3.0 if trong_lat else 1.0) | |
| n_tgt += trong_lat | |
| assert 0 < n_tgt < 50 # shard thật có cả hai phía ngưỡng | |
| def test_sample_w_di_qua_collate_va_loss(): | |
| """`sample_w` vắng → loss chạy ĐÚNG biểu thức cũ; có nhưng toàn 1.0 → | |
| số y hệt; lát target nặng hơn → loss dịch về phía cú lát target.""" | |
| def item(n, w=None): | |
| it = {"feats": np.ones((n, sn.FEAT_DIM), np.float32), | |
| "t": np.arange(n, dtype=np.float32) / 30.0, | |
| **{k: 1.0 for k in sn.LABEL_KEYS}} | |
| if w is not None: | |
| it["sample_w"] = w | |
| return it | |
| assert "sample_w" not in sn.collate([item(3), item(5)]) | |
| out = sn.collate([item(3, 1.0), item(5, 3.0)]) | |
| assert out["sample_w"].tolist() == [1.0, 3.0] | |
| b = _rand_batch() | |
| pred = {"v0_z": torch.zeros(3), "phi_vec": torch.zeros(3, 2), | |
| "ab": torch.zeros(3, 2), "ident_logit": torch.zeros(3)} | |
| l_cu = sn.shotnet_loss(pred, b, TINY) | |
| b1 = dict(b, sample_w=torch.ones(3)) | |
| l_1 = sn.shotnet_loss(pred, b1, TINY) | |
| for k in ("v0", "phi", "ab", "ident", "total"): | |
| assert l_1[k].item() == pytest.approx(l_cu[k].item(), rel=1e-6) | |
| # trọng số 3 dồn về cú 0 (label_v0=1.0, log→0 nên loss v0 nhỏ nhất) | |
| b3 = dict(b, sample_w=torch.tensor([3.0, 1.0, 1.0])) | |
| assert sn.shotnet_loss(pred, b3, TINY)["v0"].item() < l_cu["v0"].item() | |
| def test_gate_metrics_perfect_and_flipped(): | |
| labels = {"label_v0": np.array([2.0, 4.0]), | |
| "label_phi": np.array([10.0, 350.0]), | |
| "label_a": np.array([0.3, -0.3]), | |
| "label_b": np.array([0.3, -0.3]), | |
| "identifiable": np.array([1.0, 1.0])} | |
| perfect = {"v0": labels["label_v0"].copy(), | |
| "phi_deg": labels["label_phi"].copy(), | |
| "a": labels["label_a"].copy(), "b": labels["label_b"].copy(), | |
| "p_ident": np.array([0.9, 0.9])} | |
| m = sn.gate_metrics(perfect, labels) | |
| assert m["dphi_med"] == 0.0 and m["v0_relerr_med"] == 0.0 | |
| assert m["side_acc"] == 1.0 and m["vert_acc"] == 1.0 | |
| assert m["n_side"] == 2 and m["n_ident"] == 2 | |
| flipped = dict(perfect, a=-labels["label_a"], b=np.array([0.0, 0.0])) | |
| m2 = sn.gate_metrics(flipped, labels) | |
| assert m2["side_acc"] == 0.0 | |
| assert m2["vert_acc"] == 0.0 # b̂=0 → stun ≠ follow/draw GT | |