Spaces:
Sleeping
Sleeping
File size: 38,795 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 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 | # -*- coding: utf-8 -*-
"""ShotNet — net suy ngược thông số cú đánh từ quỹ đạo bẩn (BG27, design §6).
Bài toán ngược của action space RL: chuỗi quan sát {vị trí các bi (toạ độ
bàn), visibility mask, timestamp} của MỘT cú → (V0, φ, a, b) đúng 4 chiều
``cue.set_state``. Train THUẦN synthetic (``datasets/bb9_synth``, nhiễu P0
là hợp đồng — gen_synth_shots.py); loader ở đây đọc npz KHÔNG cần pooltool
để chạy được trong venv CV (thiết kế BG26).
Kiến trúc theo design §6: transformer encoder nhỏ, một **shot token** học
được đứng đầu chuỗi đọc ra output, **RoPE theo timestamp thật** (chịu fps
lệch 25–60 + frame thiếu + frame trùng — RoPE quay q/k theo GIÂY, không theo
chỉ số frame). Token = frame; feature mỗi frame là 10 slot bi cố định theo
số bi (slot 0 = cue, slot k = bi k) × (x, y, vis) + ``img_diff`` — cột cuối
là yêu cầu từ HANDOFF 23 Bất ngờ 1: frame trùng 25→30 upconvert có PTS hợp
lệ, CHỈ img_diff lộ, nên timestamp một mình không đủ cho net phân biệt.
Output head (quyết định Cowork từ HANDOFF 26b, BRIEF 27):
- ``v0`` — thước GẬY, so ``raw`` với label_v0. Dự đoán trong LOG-space
(Huber trên log V0): gate là sai số TƯƠNG ĐỐI nên residual log ∝ rel err,
Huber-raw sẽ dồn gradient về cú nhanh 8 m/s mà bỏ rơi cú 0.5 m/s.
- ``phi`` — encode (cos, sin), loss MSE tới vector đơn vị của label: liên
tục quanh 0°/360°, không cần xử wrap; φ đọc ra bằng atan2.
- ``a, b`` — Huber, **mask trên cú non-identifiable** (spin chỉ lộ qua va
chạm — design §3.2; đừng bắt net học điều không xác định, văn hoá null).
- ``identifiable`` — BCE logit.
- aux tuỳ config: ``v0_ball``/``phi_ball`` (label có sẵn trong dataset —
BG26b Bất ngờ 4); gate vẫn chấm trên V0 gậy.
Val split cắt ~2% từ TRAIN theo seed (``val_split_indices``) — 5k held-out
là TEST, cấm đụng khi train/tune (BRIEF 27 bước 1.1).
BG29 thêm HAI thứ vào **đường train** (và chỉ đường train — xem
``ShardDataset.set_train_mode``): ``SlotDropoutConfig`` làm thưa detections
bi mục tiêu on-the-fly để phủ vùng mật độ của clip thật, và
``TargetReweight`` kéo trọng số loss về lát chạm đầu ≥0.3s. Kiến trúc, thước
và mapping KHÔNG đổi một dòng.
Hằng phân lớp spin (A_SIDE_MIN/B_STUN_MAX) chép từ
``scripts/broadcast/eval_baseline_synth.py`` — GT của gate chấm ở đó; ở đây
chỉ dùng cho metric val + map (a,b) dự đoán → lớp khi cắm vào harness.
"""
from __future__ import annotations
import math
from dataclasses import dataclass, field
from pathlib import Path
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
# ------------------------------------------------------------ bảng hằng
# Bàn synth = pooltool default (spec.json bb9_synth) — KHÔNG phải bàn giải
# 1.27×2.54 của broadcast.py.
TABLE_W_M = 0.9906
TABLE_L_M = 1.9812
N_SLOTS = 10 # slot bi cố định: 0 = cue, 1..9 = bi số
FEAT_DIM = 3 * N_SLOTS + 1 # (x, y, vis) × 10 + img_diff = 31
FEAT_DIM_DELTAS = 5 * N_SLOTS + 1
# + (dx, dy) × 10 (chế độ deltas, config c2):
# φ/V0 là hàm của VẬN TỐC — c1 bắt encoder tự
# dựng chuyển động từ toạ độ tuyệt đối qua
# attention và trần ở 11.7°; đưa thẳng finite
# difference vào layer 1
IMG_DIFF_CLIP = 2.0 # img_diff cú nhanh ~0.8–2; frame đầu sentinel
# −1 → clip về 0 (t=0 tự khai frame đầu)
# GT spin — CHÉP từ eval_baseline_synth.py (một nguồn, hai bản có chủ đích:
# src/ không import được scripts/; test khoá hai bản bằng nhau nếu cần):
B_STUN_MAX = 0.10 # |b| ≤ 0.10 = stun
A_SIDE_MIN = 0.10 # |a| ≤ 0.10 = neutral (side không chấm)
# ------------------------------------------------------------- featurize
def featurize_shot(shot: dict, w: float = TABLE_W_M, l: float = TABLE_L_M,
deltas: bool = False) -> tuple[np.ndarray, np.ndarray]:
"""Một cú (dict của ``iter_shots``/``ShardDataset``) → (feats (F, 31
hoặc 51 nếu ``deltas``), t (F,)) float32.
Chuẩn hoá vị trí ĐẲNG HƯỚNG quanh tâm bàn: x' = (x − w/2)·2/l,
y' = (y − l/2)·2/l — khoảng cách vật lý giữ nguyên tỷ lệ hai trục (net
ước lượng vận tốc từ hiệu vị trí). Bi không thấy (covered=False) hoặc
slot không có bi trên bàn: (0, 0, vis=0).
``deltas`` (config c2): thêm khối (dx, dy) per-slot = hiệu vị trí với
frame TRƯỚC, cùng thang chuẩn hoá, chỉ khi CẢ HAI frame thấy bi (khác
đi bằng 0 — không nhảy vọt qua gap); frame trùng upconvert cho Δ=0,
kết hợp img_diff net phân biệt được dup với đứng yên. img_diff luôn là
cột CUỐI ở cả hai chế độ.
"""
xy, cov = shot["xy"], shot["covered"]
ball_ids = np.asarray(shot["ball_ids"], dtype=np.int64)
F_n = xy.shape[0]
dim = FEAT_DIM_DELTAS if deltas else FEAT_DIM
feats = np.zeros((F_n, dim), dtype=np.float32)
xn = ((xy[:, :, 0] - w / 2.0) * (2.0 / l)).astype(np.float32)
yn = ((xy[:, :, 1] - l / 2.0) * (2.0 / l)).astype(np.float32)
vis = cov.astype(np.float32)
for j, s in enumerate(ball_ids):
feats[:, 3 * s] = xn[:, j] * vis[:, j]
feats[:, 3 * s + 1] = yn[:, j] * vis[:, j]
feats[:, 3 * s + 2] = vis[:, j]
if deltas:
both = np.zeros(F_n, dtype=np.float32)
both[1:] = vis[1:, j] * vis[:-1, j]
d0 = 3 * N_SLOTS + 2 * s
feats[1:, d0] = (xn[1:, j] - xn[:-1, j]) * both[1:]
feats[1:, d0 + 1] = (yn[1:, j] - yn[:-1, j]) * both[1:]
feats[:, -1] = np.clip(shot["img_diff"], 0.0, IMG_DIFF_CLIP)
return feats, np.asarray(shot["t"], dtype=np.float32)
LABEL_KEYS = ("label_v0", "label_phi", "label_a", "label_b", "identifiable",
"v0_ball", "phi_ball", "fps", "upconvert", "scratch")
# ------------------------- augmentation "detections thưa bi mục tiêu" (BG29)
# Chẩn đoán BG28 (ĐO ĐƯỢC, không đoán): net c3 sập trên clip thật vì PHÂN
# PHỐI mật độ detection bi mục tiêu, không phải vì kiến trúc hay input path
# — synth train 95–96% visibility (≈5.3 det không-cue/frame) trong khi cú 11
# thật chỉ 1.01 det/frame, cú 12 3.09. Augmentation dưới đây làm THƯA
# detections ngay trong loader train (on-the-fly, KHÔNG sinh dataset mới ra
# đĩa — BRIEF 29 bước 1.1) để dải huấn luyện phủ được vùng thật đã đo.
#
# Dropout ra mask 0 ĐÚNG như bi mất frame (featurize_shot nhân toạ độ với
# vis, khối delta nhân tích vis hai frame) — không đụng img_diff/timestamp,
# không đụng feature nào khác.
@dataclass
class SlotDropoutConfig:
"""Dropout slot bi mục tiêu, mỗi cú một keep-rate.
``p_apply``: tỉ lệ cú ĐƯỢC augment; phần còn lại đi qua loader **bit
giống** đường cũ (giữ nguyên vùng 95–96% visibility của synth trong
phân phối train — đó cũng là đầu trên của dải).
``keep_min``/``keep_max``: keep-rate ~ U(min, max), Bernoulli độc lập
trên từng ô (frame, slot không-cue). Dải mặc định [0.10, 0.95] cộng với
nhánh không-augment phủ trọn vùng thật đã đo: cú 11 ~15% visibility
per-ball, cú 12 ~50%, synth gốc 95–96%.
``cue`` KHÔNG bao giờ bị dropout thêm: track cue trên clip thật vốn tốt
(coverage 0.90–0.97 ở BG28) và nhiễu đầu cú đã nằm trong hợp đồng P0.
"""
p_apply: float = 0.75
keep_min: float = 0.10
keep_max: float = 0.95
seed: int = 20260812
def apply_slot_dropout(shot: dict, cfg: SlotDropoutConfig,
rng: np.random.Generator) -> tuple[dict, float]:
"""Thưa detections bi mục tiêu của MỘT cú → (shot mới, keep-rate đã
dùng; 1.0 = không augment).
KHÔNG sửa tại chỗ: ``covered`` của ``ShardDataset.raw_shot`` là view vào
mảng shard nạp sẵn trong RAM — ghi đè lên đó là hỏng dataset cho mọi
epoch sau.
"""
cov = shot["covered"]
tgt = np.flatnonzero(np.asarray(shot["ball_ids"]) != 0)
if tgt.size == 0 or rng.random() >= cfg.p_apply:
return shot, 1.0
keep = float(rng.uniform(cfg.keep_min, cfg.keep_max))
cov = np.array(cov, dtype=bool, copy=True)
cov[:, tgt] &= rng.random((cov.shape[0], tgt.size)) < keep
out = dict(shot)
out["covered"] = cov
return out, keep
# Lát target-matched — CHÉP từ scripts/broadcast/eval_baseline_synth.py
# (T_FIRST_TARGET_S; src/ không import được scripts/, nếp B_STUN_MAX ở trên).
# Quần thể khớp clip thật P0 và là thước chính thức của gate P2′.
T_FIRST_TARGET_S = 0.3
@dataclass
class TargetReweight:
"""Reweight về lát target (BRIEF 29 bước 1.2): cú có chạm đầu ≥
``t_first_s`` được nhân trọng số ``weight`` trong loss.
Chọn **weight loss** thay vì oversample có chủ đích: oversample lát
13.6% với hệ số 3 sẽ kéo dài epoch ~27%, phá điều kiện "ngân sách bước
y hệt c3" của BRIEF 29 bước 2.1 — hai config sẽ không so được với c3 nữa.
Trọng số áp cho MỌI head (v0/phi/ab/ident) qua trung bình có trọng số,
nên thang gradient tổng giữ nguyên ~1 (bài học c1 BG27: lệch thang
gradient là head chết).
"""
weight: float = 3.0
t_first_s: float = T_FIRST_TARGET_S
def t_first_contact_s(shot: dict) -> float:
"""Chạm đầu của cue kể từ strike: min(t_first_bb, t_first_cush) — NaN
nếu cú không có va chạm nào. Cùng công thức eval_baseline_synth."""
vals = [float(v) for v in (shot["t_first_bb"], shot["t_first_cush"])
if not math.isnan(float(v))]
return min(vals) if vals else math.nan
# ------------------------------------- sim→real: dựng shot từ track app (BG28)
# Khe input sim→real là rủi ro chính của bàn giao 28: net train với slot bi
# mục tiêu bền theo thời gian (slot k = bi k trong sim), còn pipeline app chỉ
# track cue — detections mỗi frame CÓ các bi khác nhưng chưa nối. Dựng slot
# bằng greedy nearest-neighbor qua frame (BRIEF 28 bước 2.2): KHÔNG cần
# BallID — slot chỉ cần BỀN theo thời gian như slot sim; bi mất frame →
# mask 0. Đơn giản có chủ đích.
#
# Hằng nối det→slot MƯỢN đúng số đã nghiệm thu của track cue (broadcast.py
# VMAX_MPS/GATE_SLACK_M/BALL_DEDUP_M — không import được vì src/ không kéo
# scripts/, và broadcast.py là tầng cấm sửa; chép số kèm nguồn, nếp
# B_STUN_MAX ở trên):
SLOT_GATE_MPS = 11.0 # trần vận tốc nối det giữa 2 frame (= VMAX_MPS)
SLOT_GATE_SLACK_M = 0.06 # nhiễu tâm bbox (~2R) (= GATE_SLACK_M)
SLOT_DEDUP_M = 0.03 # 2 det cùng frame < ~1R = double-detect, bỏ det
# sau (= BALL_DEDUP_M)
MAX_TARGET_SLOTS = N_SLOTS - 1 # tối đa 9 slot bi mục tiêu (slot 0 = cue)
def build_target_slots(times: np.ndarray,
dets_by_frame: list[list[tuple[float, float]]]
) -> tuple[np.ndarray, np.ndarray]:
"""Nối detections KHÔNG-cue rời rạc thành slot bền theo thời gian.
``times`` (F,) giây; ``dets_by_frame`` list F phần tử, mỗi phần tử là
list (x, y) mét của frame đó. Trả (slot_xy (F, K, 2), slot_vis (F, K))
với K ≤ MAX_TARGET_SLOTS slot theo thứ tự XUẤT HIỆN.
Mỗi frame: dedup double-detect (< SLOT_DEDUP_M) rồi ghép det↔slot greedy
theo khoảng cách tăng dần, gate = SLOT_GATE_MPS·Δt_kể_từ_lần_thấy_cuối +
SLOT_GATE_SLACK_M (bi có thể đã chạy trong lúc mất frame — gate nở theo
thời gian mất). Det không ghép được slot nào → slot MỚI (quá
MAX_TARGET_SLOTS thì bỏ det). Frame không có det của slot → vis 0, đúng
ngữ nghĩa ``covered`` của loader synth.
"""
F = len(times)
slots: list[dict] = [] # {x, y, t_last}
raw_xy: list[np.ndarray] = [] # (F, 2) per slot — cấp phát lười
raw_vis: list[np.ndarray] = []
for f in range(F):
t = float(times[f])
dets: list[tuple[float, float]] = []
for p in dets_by_frame[f]:
if any(math.hypot(p[0] - q[0], p[1] - q[1]) < SLOT_DEDUP_M
for q in dets):
continue
dets.append((float(p[0]), float(p[1])))
if not dets:
continue
# mọi cặp (slot, det) trong gate, xét theo khoảng cách tăng dần
pairs = []
for si, s in enumerate(slots):
gate = SLOT_GATE_MPS * max(t - s["t_last"], 0.0) \
+ SLOT_GATE_SLACK_M
for di, p in enumerate(dets):
d = math.hypot(p[0] - s["x"], p[1] - s["y"])
if d <= gate:
pairs.append((d, si, di))
pairs.sort(key=lambda z: z[0])
used_s: set[int] = set()
used_d: set[int] = set()
for d, si, di in pairs:
if si in used_s or di in used_d:
continue
used_s.add(si)
used_d.add(di)
s = slots[si]
s["x"], s["y"], s["t_last"] = dets[di][0], dets[di][1], t
raw_xy[si][f] = dets[di]
raw_vis[si][f] = True
for di, p in enumerate(dets):
if di in used_d or len(slots) >= MAX_TARGET_SLOTS:
continue
slots.append({"x": p[0], "y": p[1], "t_last": t})
raw_xy.append(np.zeros((F, 2), dtype=np.float32))
raw_vis.append(np.zeros(F, dtype=bool))
raw_xy[-1][f] = p
raw_vis[-1][f] = True
if not slots:
return (np.zeros((F, 0, 2), dtype=np.float32),
np.zeros((F, 0), dtype=bool))
return (np.stack(raw_xy, axis=1).astype(np.float32),
np.stack(raw_vis, axis=1))
def shot_from_track(rows: list[dict], others: list[dict] | None
) -> dict:
"""Dựng dict cú (cùng shape ``iter_shots``/``featurize_shot`` cần) từ
ĐÚNG input của ``broadcast.analyze_track``: ``rows`` per-frame của track
cue (t_s, covered, table_x_m/y_m, img_diff) + ``others`` detection
KHÔNG-cue rời rạc ({t_s, x_m, y_m}). Một nguồn sự thật với loader —
caller featurize bằng ``featurize_shot(shot, w=<bàn thật>, l=<bàn
thật>)``; chuẩn hoá đẳng hướng của featurize làm phần còn lại (hai bàn
cùng tỷ lệ 2:1 nên toạ độ chuẩn hoá trùng khớp phân phối synth).
KHÔNG dedup frame trùng, KHÔNG smooth — net train trên chuỗi thô có
frame dup + img_diff, đưa thô vào đúng như loader (HANDOFF 23 BN1).
Timestamp trừ mốc frame đầu (synth luôn bắt đầu ở 0; PTS container có
thể có offset).
"""
F = len(rows)
t0 = float(rows[0]["t_s"])
t = np.array([float(r["t_s"]) - t0 for r in rows], dtype=np.float32)
cue_xy = np.zeros((F, 2), dtype=np.float32)
cue_vis = np.zeros(F, dtype=bool)
img_diff = np.zeros(F, dtype=np.float32)
for i, r in enumerate(rows):
on = str(r.get("covered", "0")) in ("1", "True", "true")
cue_vis[i] = on
if on:
cue_xy[i] = (float(r["table_x_m"]), float(r["table_y_m"]))
img_diff[i] = float(r.get("img_diff", -1.0))
# gom detection theo frame: t_s của others là ĐÚNG t_s của frame sinh ra
# nó (analyze_clip đọc từ cùng meta), map bằng giá trị
by_t: dict[float, list[tuple[float, float]]] = {}
for o in (others or []):
by_t.setdefault(float(o["t_s"]), []).append(
(float(o["x_m"]), float(o["y_m"])))
dets_by_frame = [by_t.get(float(r["t_s"]), []) for r in rows]
slot_xy, slot_vis = build_target_slots(t, dets_by_frame)
K = slot_xy.shape[1]
xy = np.concatenate([cue_xy[:, None, :], slot_xy], axis=1)
covered = np.concatenate([cue_vis[:, None], slot_vis], axis=1)
return {"t": t, "xy": xy, "covered": covered, "img_diff": img_diff,
"ball_ids": np.arange(K + 1, dtype=np.int64),
"n_frames": F, "n_balls": K + 1}
# --------------------------------------------------------------- dataset
class ShardDataset:
"""Nạp danh sách shard npz vào RAM (ragged + offset, ~2.4GB cho 150k),
trả từng cú đã featurize. KHÔNG cần pooltool; không phụ thuộc torch cho
tới ``collate`` (test loader chạy được thuần numpy).
``indices``: chỉ số TOÀN CỤC (theo thứ tự shard đã sort + local idx) —
dùng cho train/val split. ``limit_shots``: cắt mỗi shard (smoke test).
``augment``/``reweight`` (BG29): CHỈ có hiệu lực khi ``set_train_mode(
True)``. Mặc định TẮT — mọi đường eval (harness held-out, worker app,
val mỗi epoch) phải thấy dữ liệu y hệt c3, nếu không số không so được.
"""
_META = ("label_v0", "label_phi", "label_a", "label_b", "v0_ball",
"phi_ball", "identifiable", "scratch", "fps", "upconvert",
"t_first_bb", "t_first_cush", "n_frames", "n_balls", "shot_idx")
def __init__(self, shard_paths: list[Path], limit_shots: int = 0,
deltas: bool = False,
augment: SlotDropoutConfig | None = None,
reweight: TargetReweight | None = None):
self.deltas = deltas
self.augment = augment
self.reweight = reweight
self._train_mode = False
self._epoch = 0
self.shards = []
self._index: list[tuple[int, int]] = []
for si, p in enumerate(sorted(Path(q) for q in shard_paths)):
z = np.load(p)
S = len(z["shot_idx"])
if limit_shots:
S = min(S, limit_shots)
sh = {k: np.asarray(z[k]) for k in self._META}
sh["xy"] = np.asarray(z["xy"], dtype=np.float32)
sh["covered"] = np.asarray(z["covered"]).astype(bool)
sh["img_diff"] = np.asarray(z["img_diff"], dtype=np.float32)
sh["ball_ids"] = np.asarray(z["ball_ids"])
sh["xy_off"] = np.asarray(z["xy_off"], dtype=np.int64)
sh["f_off"] = np.asarray(z["f_off"], dtype=np.int64)
sh["b_off"] = np.asarray(z["b_off"], dtype=np.int64)
self.shards.append(sh)
self._index.extend((si, li) for li in range(S))
self.lengths = np.array(
[self.shards[si]["n_frames"][li] for si, li in self._index],
dtype=np.int64)
def __len__(self) -> int:
return len(self._index)
def set_train_mode(self, on: bool, epoch: int = 0) -> None:
"""Bật/tắt augmentation + reweight, chốt epoch cho RNG.
Train bật (epoch khác nhau → mask khác nhau, vẫn tái lập được từ
``[seed, epoch, i]``); val/eval TẮT. Một cờ cho cả hai vì cả hai chỉ
có nghĩa ở đường train — số val phải so được với c3 và giữa các
config.
"""
self._train_mode = bool(on)
self._epoch = int(epoch)
def raw_shot(self, i: int) -> dict:
"""Dict cú thô cùng shape với ``gen_synth_shots.iter_shots``."""
si, li = self._index[i]
sh = self.shards[si]
F_n = int(sh["n_frames"][li])
B = int(sh["n_balls"][li])
x0, f0, b0 = (int(sh["xy_off"][li]), int(sh["f_off"][li]),
int(sh["b_off"][li]))
shot = {k: sh[k][li].item() for k in self._META
if k not in ("n_frames", "n_balls")}
shot["n_frames"], shot["n_balls"] = F_n, B
shot["t"] = (np.arange(F_n) / float(sh["fps"][li])).astype(np.float32)
shot["xy"] = sh["xy"][x0:x0 + F_n * B].reshape(F_n, B, 2)
shot["covered"] = sh["covered"][x0:x0 + F_n * B].reshape(F_n, B)
shot["img_diff"] = sh["img_diff"][f0:f0 + F_n]
shot["ball_ids"] = sh["ball_ids"][b0:b0 + B]
return shot
def __getitem__(self, i: int) -> dict:
shot = self.raw_shot(i)
keep = 1.0
if self._train_mode and self.augment is not None:
# RNG khoá theo [seed, epoch, i]: cùng bộ ba → cùng mask (tái
# lập được từng epoch), khác epoch → mask khác (augmentation
# thật sự on-the-fly chứ không phải một bản thưa cố định)
rng = np.random.default_rng(
[self.augment.seed, self._epoch, int(i)])
shot, keep = apply_slot_dropout(shot, self.augment, rng)
feats, t = featurize_shot(shot, deltas=self.deltas)
item = {"feats": feats, "t": t}
for k in LABEL_KEYS:
item[k] = float(shot[k])
item["keep_rate"] = keep # chẩn đoán/test, không vào collate
if self._train_mode and self.reweight is not None:
tf = t_first_contact_s(shot)
item["sample_w"] = (self.reweight.weight
if (not math.isnan(tf))
and tf >= self.reweight.t_first_s else 1.0)
return item
def val_split_indices(n_total: int, val_frac: float,
seed: int) -> tuple[np.ndarray, np.ndarray]:
"""Split val ~``val_frac`` cắt từ TRAIN theo seed — tất định, rời nhau,
độc lập machine/shard. Trả (train_idx, val_idx) đã sort."""
perm = np.random.default_rng(seed).permutation(n_total)
n_val = int(round(n_total * val_frac))
return np.sort(perm[n_val:]), np.sort(perm[:n_val])
class BucketBatcher:
"""Gom batch theo NGÂN SÁCH TOKEN (max_len_trong_batch × cỡ_batch ≤
budget) sau khi sort-theo-độ-dài có jitter — chuỗi 13..1192 frame mà
batch cố định thì hoặc tràn VRAM ở đuôi dài hoặc phí padding ở đuôi
ngắn (VRAM 3070 = 8GB, BRIEF bối cảnh 6). Thứ tự batch xáo theo
[seed, epoch] — tái lập được từng epoch."""
def __init__(self, lengths: np.ndarray, indices: np.ndarray,
token_budget: int = 40_000, max_batch: int = 256,
seed: int = 0):
self.lengths, self.indices = lengths, np.asarray(indices)
self.token_budget, self.max_batch = token_budget, max_batch
self.seed = seed
def epoch_batches(self, epoch: int, shuffle: bool = True) -> list[np.ndarray]:
rng = np.random.default_rng([self.seed, epoch])
lens = self.lengths[self.indices].astype(np.float64)
if shuffle:
lens = lens * rng.uniform(0.9, 1.1, len(lens))
order = self.indices[np.argsort(lens, kind="stable")]
batches, cur, cur_max = [], [], 0
for i in order:
li = int(self.lengths[i])
new_max = max(cur_max, li)
if cur and ((len(cur) + 1) * new_max > self.token_budget
or len(cur) >= self.max_batch):
batches.append(np.array(cur))
cur, cur_max = [], 0
new_max = li
cur.append(int(i))
cur_max = new_max
if cur:
batches.append(np.array(cur))
if shuffle:
rng.shuffle(batches)
return batches
def collate(items: list[dict]) -> dict:
"""Pad về max-len trong batch → tensor torch. ``mask`` True = frame
thật, False = padding.
``sample_w`` (BG29 reweight) CHỈ xuất hiện khi loader phát ra nó —
batch không có khoá này đi qua ĐÚNG code path loss cũ (c3/c4a), không
phải nhánh trung bình-có-trọng-số.
"""
B = len(items)
T = max(len(it["t"]) for it in items)
dim = items[0]["feats"].shape[1]
x = torch.zeros(B, T, dim, dtype=torch.float32)
t = torch.zeros(B, T, dtype=torch.float32)
mask = torch.zeros(B, T, dtype=torch.bool)
for i, it in enumerate(items):
n = len(it["t"])
x[i, :n] = torch.from_numpy(np.ascontiguousarray(it["feats"]))
t[i, :n] = torch.from_numpy(np.ascontiguousarray(it["t"]))
mask[i, :n] = True
out = {"x": x, "t": t, "mask": mask}
for k in LABEL_KEYS:
out[k] = torch.tensor([it[k] for it in items], dtype=torch.float32)
if any("sample_w" in it for it in items):
out["sample_w"] = torch.tensor(
[float(it.get("sample_w", 1.0)) for it in items],
dtype=torch.float32)
return out
# ----------------------------------------------------------------- model
@dataclass
class ShotNetConfig:
d_model: int = 128
n_layers: int = 4
n_heads: int = 4
d_ff: int = 512
dropout: float = 0.1
feat_dim: int = FEAT_DIM
use_deltas: bool = False # featurize thêm khối (dx, dy) — xem
# featurize_shot; feat_dim tự nâng 51
ab_scale: float = 1.0 # target (a,b)/ab_scale khi tính loss —
# 0.4 (c2) đưa target về ±1 để gradient
# spin ngang phi (c1: head a sập về hằng
# 0, corr 0.02 — gradient 0.05 quá yếu)
# RoPE theo GIÂY: chu kỳ hình học từ min (phân giải ~nửa bước frame 60fps
# ×2) tới max (dài hơn clip dài nhất ~21.5s)
rope_min_period_s: float = 0.04
rope_max_period_s: float = 64.0
shot_token_t_s: float = -1.0 # timestamp gán cho shot token (trước clip)
v0_log: bool = True # Huber trên log V0 (xem docstring module)
huber_v0: float = 0.25
huber_ab: float = 0.05
phi_loss: str = "vec_mse" # "vec_mse" (c1/c2) | "ang_huber" (c3):
# vec-MSE có gradient ∝ sin(Δ) nên cú vô
# vọng (Δ~90°) thống trị; gate chấm MEDIAN
# — Huber trên GÓC (linear ngoài delta) ép
# khối giữa phân phối về dưới bar
phi_huber_deg: float = 5.0 # delta của ang_huber (độ)
w_v0: float = 1.0
w_phi: float = 2.0
w_ab: float = 1.0
w_ident: float = 0.2
aux_ball: bool = False # head phụ v0_ball/phi_ball (BG26b BN4)
w_aux: float = 0.5
def __post_init__(self):
if self.use_deltas and self.feat_dim == FEAT_DIM:
self.feat_dim = FEAT_DIM_DELTAS
class RoPE(nn.Module):
"""Rotary embedding theo timestamp THỰC (giây), tần số hình học."""
def __init__(self, head_dim: int, min_period_s: float,
max_period_s: float):
super().__init__()
n_pairs = head_dim // 2
k = torch.arange(n_pairs, dtype=torch.float32) / max(n_pairs - 1, 1)
omega = (2.0 * math.pi / max_period_s) * \
(max_period_s / min_period_s) ** k
self.register_buffer("omega", omega, persistent=False)
def rotate(self, q: torch.Tensor, t: torch.Tensor) -> torch.Tensor:
# q: (B, H, T, Dh) · t: (B, T) → quay từng cặp (2i, 2i+1)
ang = t[:, None, :, None] * self.omega[None, None, None, :]
cos, sin = torch.cos(ang), torch.sin(ang)
q1, q2 = q[..., 0::2], q[..., 1::2]
out = torch.empty_like(q)
out[..., 0::2] = q1 * cos - q2 * sin
out[..., 1::2] = q1 * sin + q2 * cos
return out
class EncoderLayer(nn.Module):
def __init__(self, cfg: ShotNetConfig, rope: RoPE):
super().__init__()
self.h = cfg.n_heads
self.dh = cfg.d_model // cfg.n_heads
self.rope = rope
self.ln1 = nn.LayerNorm(cfg.d_model)
self.qkv = nn.Linear(cfg.d_model, 3 * cfg.d_model)
self.proj = nn.Linear(cfg.d_model, cfg.d_model)
self.ln2 = nn.LayerNorm(cfg.d_model)
self.ff = nn.Sequential(
nn.Linear(cfg.d_model, cfg.d_ff), nn.GELU(),
nn.Dropout(cfg.dropout), nn.Linear(cfg.d_ff, cfg.d_model))
self.drop = nn.Dropout(cfg.dropout)
def forward(self, h: torch.Tensor, t: torch.Tensor,
key_mask: torch.Tensor) -> torch.Tensor:
B, T, D = h.shape
q, k, v = self.qkv(self.ln1(h)).chunk(3, dim=-1)
q = q.view(B, T, self.h, self.dh).transpose(1, 2)
k = k.view(B, T, self.h, self.dh).transpose(1, 2)
v = v.view(B, T, self.h, self.dh).transpose(1, 2)
q, k = self.rope.rotate(q, t), self.rope.rotate(k, t)
att = F.scaled_dot_product_attention(
q, k, v, attn_mask=key_mask[:, None, None, :])
att = att.transpose(1, 2).reshape(B, T, D)
h = h + self.drop(self.proj(att))
h = h + self.drop(self.ff(self.ln2(h)))
return h
class ShotNet(nn.Module):
def __init__(self, cfg: ShotNetConfig):
super().__init__()
self.cfg = cfg
self.inp = nn.Linear(cfg.feat_dim, cfg.d_model)
self.shot_token = nn.Parameter(torch.zeros(1, 1, cfg.d_model))
nn.init.normal_(self.shot_token, std=0.02)
rope = RoPE(cfg.d_model // cfg.n_heads, cfg.rope_min_period_s,
cfg.rope_max_period_s)
self.layers = nn.ModuleList(
EncoderLayer(cfg, rope) for _ in range(cfg.n_layers))
self.ln_f = nn.LayerNorm(cfg.d_model)
d = cfg.d_model
self.head_v0 = nn.Linear(d, 1)
self.head_phi = nn.Linear(d, 2) # (cos, sin)
self.head_ab = nn.Linear(d, 2) # (a, b)
self.head_ident = nn.Linear(d, 1)
if cfg.aux_ball:
self.head_v0_ball = nn.Linear(d, 1)
self.head_phi_ball = nn.Linear(d, 2)
# bias V0 khởi tạo giữa dải log [0.5, 8] — đỡ vài epoch đầu
nn.init.constant_(self.head_v0.bias, math.log(2.0))
def forward(self, x: torch.Tensor, t: torch.Tensor,
mask: torch.Tensor) -> dict:
B = x.shape[0]
h = self.inp(x)
tok = self.shot_token.expand(B, 1, -1)
h = torch.cat([tok, h], dim=1)
t_full = torch.cat(
[torch.full((B, 1), self.cfg.shot_token_t_s,
dtype=t.dtype, device=t.device), t], dim=1)
key_mask = torch.cat(
[torch.ones(B, 1, dtype=torch.bool, device=mask.device), mask],
dim=1)
for layer in self.layers:
h = layer(h, t_full, key_mask)
g = self.ln_f(h[:, 0])
out = {"v0_z": self.head_v0(g).squeeze(-1),
"phi_vec": self.head_phi(g),
"ab": self.head_ab(g),
"ident_logit": self.head_ident(g).squeeze(-1)}
if self.cfg.aux_ball:
out["v0_ball_z"] = self.head_v0_ball(g).squeeze(-1)
out["phi_ball_vec"] = self.head_phi_ball(g)
return out
@torch.no_grad()
def predict(self, x: torch.Tensor, t: torch.Tensor,
mask: torch.Tensor) -> dict:
"""Suy đoán → numpy: v0 (m/s, thước gậy), phi_deg [0,360), a, b,
p_ident."""
self.eval()
out = self.forward(x, t, mask)
v0 = torch.exp(out["v0_z"]) if self.cfg.v0_log else out["v0_z"]
phi = torch.rad2deg(torch.atan2(out["phi_vec"][:, 1],
out["phi_vec"][:, 0])) % 360.0
return {"v0": v0.cpu().numpy(),
"phi_deg": phi.cpu().numpy(),
"a": (out["ab"][:, 0] * self.cfg.ab_scale).cpu().numpy(),
"b": (out["ab"][:, 1] * self.cfg.ab_scale).cpu().numpy(),
"p_ident": torch.sigmoid(out["ident_logit"]).cpu().numpy()}
# ------------------------------------------------------------------ loss
def _phi_vec_loss(pred_vec: torch.Tensor,
phi_deg: torch.Tensor) -> torch.Tensor:
"""MSE tới vector đơn vị (cos φ, sin φ) — liên tục quanh 0°/360°."""
rad = torch.deg2rad(phi_deg)
target = torch.stack([torch.cos(rad), torch.sin(rad)], dim=-1)
return F.mse_loss(pred_vec, target, reduction="none").sum(-1)
def _phi_ang_huber(pred_vec: torch.Tensor, phi_deg: torch.Tensor,
delta_deg: float) -> torch.Tensor:
"""Huber trên KHOẢNG CÁCH GÓC (radian, delta cho bằng độ): Δ qua
atan2(sin, cos) của hiệu góc — liên tục quanh 0°/360°. So vec-MSE:
cùng độ cong ~Δ² khi nhỏ, nhưng ngoài delta gradient KHÔNG tăng theo
sin(Δ) — cú vô vọng (Δ~90°) bị chặn trần, khối giữa phân phối (thứ
quyết định MEDIAN của gate) không bị đuôi đè. Cộng phạt nhẹ giữ
‖pred_vec‖ ≈ 1 (hướng là tín hiệu, norm không được sập)."""
rad = torch.deg2rad(phi_deg)
norm = torch.linalg.vector_norm(pred_vec, dim=-1)
u = pred_vec / (norm[..., None] + 1e-6)
cos_d = u[..., 0] * torch.cos(rad) + u[..., 1] * torch.sin(rad)
sin_d = u[..., 1] * torch.cos(rad) - u[..., 0] * torch.sin(rad)
d_rad = torch.atan2(sin_d, cos_d).abs()
hub = F.huber_loss(d_rad, torch.zeros_like(d_rad),
delta=math.radians(delta_deg), reduction="none")
return hub + 0.01 * (norm - 1.0) ** 2
def shotnet_loss(out: dict, batch: dict, cfg: ShotNetConfig) -> dict:
"""Loss tổng + từng head. (a, b) chỉ tính trên cú identifiable.
``batch["sample_w"]`` (BG29 bước 1.2) — nếu CÓ, mọi head thu về bằng
TRUNG BÌNH CÓ TRỌNG SỐ ``Σwl/Σw`` thay vì trung bình thường: thang loss
(và do đó thang gradient) giữ nguyên ~1 dù trọng số lát target là 3.
KHÔNG có khoá này → chạy đúng biểu thức cũ, không đổi một phép tính
(c3/c4a phải so được từng số).
"""
w = batch.get("sample_w")
if w is None:
def red(x):
return x.mean()
def red_m(x, m):
return (x * m).sum() / m.sum().clamp(min=1.0)
else:
def red(x):
return (x * w).sum() / w.sum().clamp(min=1e-6)
def red_m(x, m):
return (x * m * w).sum() / (m * w).sum().clamp(min=1e-6)
v0_t = batch["label_v0"]
v0_target = torch.log(v0_t) if cfg.v0_log else v0_t
l_v0 = red(F.huber_loss(out["v0_z"], v0_target, delta=cfg.huber_v0,
reduction="none"))
if cfg.phi_loss == "ang_huber":
l_phi = red(_phi_ang_huber(out["phi_vec"], batch["label_phi"],
cfg.phi_huber_deg))
else:
l_phi = red(_phi_vec_loss(out["phi_vec"], batch["label_phi"]))
ab_target = torch.stack([batch["label_a"], batch["label_b"]],
dim=-1) / cfg.ab_scale
l_ab_per = F.huber_loss(out["ab"], ab_target, delta=cfg.huber_ab,
reduction="none").mean(-1)
m = batch["identifiable"]
l_ab = red_m(l_ab_per, m)
l_id = red(F.binary_cross_entropy_with_logits(out["ident_logit"], m,
reduction="none"))
total = (cfg.w_v0 * l_v0 + cfg.w_phi * l_phi + cfg.w_ab * l_ab
+ cfg.w_ident * l_id)
parts = {"v0": l_v0, "phi": l_phi, "ab": l_ab, "ident": l_id}
if cfg.aux_ball:
vb = batch["v0_ball"].clamp(min=1e-3)
l_avb = red(F.huber_loss(out["v0_ball_z"],
torch.log(vb) if cfg.v0_log else vb,
delta=cfg.huber_v0, reduction="none"))
if cfg.phi_loss == "ang_huber":
l_apb = red(_phi_ang_huber(out["phi_ball_vec"],
batch["phi_ball"], cfg.phi_huber_deg))
else:
l_apb = red(_phi_vec_loss(out["phi_ball_vec"],
batch["phi_ball"]))
total = total + cfg.w_aux * (l_avb + l_apb)
parts.update({"aux_v0_ball": l_avb, "aux_phi_ball": l_apb})
parts["total"] = total
return parts
# --------------------------------------------------------------- metrics
def circ_diff_deg(x: np.ndarray, y: np.ndarray) -> np.ndarray:
"""|Δφ| trên vòng tròn, độ — cùng công thức eval_baseline_synth."""
return np.abs(((np.asarray(x) - np.asarray(y) + 180.0) % 360.0) - 180.0)
def spin_classes(a: np.ndarray, b: np.ndarray) -> tuple[np.ndarray, np.ndarray]:
"""(a, b) → (vert, side) theo đúng ngưỡng GT của harness. vert luôn có
lớp (follow/draw/stun); side rỗng khi |a| ≤ A_SIDE_MIN (neutral)."""
a, b = np.asarray(a), np.asarray(b)
vert = np.where(b > B_STUN_MAX, "follow",
np.where(b < -B_STUN_MAX, "draw", "stun"))
side = np.where(a > A_SIDE_MIN, "side-L",
np.where(a < -A_SIDE_MIN, "side-R", ""))
return vert, side
def gate_metrics(pred: dict, labels: dict) -> dict:
"""4 chỉ số gate G-27.3 đo trên một tập (val mỗi epoch):
- dphi_med: median |Δφ| toàn tập (độ)
- v0_relerr_med: median |V0̂−V0|/V0 (raw, thước gậy) toàn tập
- side_acc: dấu a (sign, KHÔNG dead-zone — net không abstain)
trên cú identifiable có GT side ≠ neutral
- vert_acc: 3 lớp follow/stun/draw (quantize b̂ theo B_STUN_MAX)
trên TOÀN identifiable
"""
ident = np.asarray(labels["identifiable"]) > 0.5
dphi = circ_diff_deg(pred["phi_deg"], labels["label_phi"])
relerr = np.abs(pred["v0"] - labels["label_v0"]) / labels["label_v0"]
gt_vert, gt_side = spin_classes(labels["label_a"], labels["label_b"])
pv, _ = spin_classes(pred["a"], pred["b"])
side_mask = ident & (gt_side != "")
side_ok = (np.sign(pred["a"]) == np.sign(labels["label_a"]))[side_mask]
vert_ok = (pv == gt_vert)[ident]
id_acc = ((pred["p_ident"] > 0.5) == ident).mean()
return {"dphi_med": float(np.median(dphi)),
"v0_relerr_med": float(np.median(relerr)),
"side_acc": float(side_ok.mean()) if side_ok.size else float("nan"),
"vert_acc": float(vert_ok.mean()) if vert_ok.size else float("nan"),
"n_side": int(side_mask.sum()), "n_ident": int(ident.sum()),
"ident_acc": float(id_acc)}
|