poolcoach / tests /test_shotnet.py
masterdanh's picture
deploy: snapshot for HF Space
78738de
Raw
History Blame Contribute Delete
28.5 kB
# -*- 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
@pytest.fixture(scope="module")
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