| |
| """Build causal raw-RGB frame caches for L1 situation understanding. |
| |
| The cache aligns policy decision rows to processed fight videos: |
| |
| policy video_t -> latent cell k = round((video_t - 0.125) / 0.25) |
| history cell k + offset -> video frame 4 * history_cell + frame_in_cell |
| |
| By default only strictly negative offsets are accepted, so no target/onset |
| frame is used. Use --allow_current_frame only for a clearly labeled |
| current-observation ablation. |
| Frames are stored as resized uint8 RGB tensors to make training runs |
| repeatable without repeatedly seeking through mp4 files. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import argparse |
| import json |
| import os |
| import sys |
| from collections import Counter, defaultdict |
| from typing import Any, Dict, List, Sequence, Tuple |
|
|
| import multiprocessing as mp |
| from concurrent.futures import ProcessPoolExecutor, as_completed |
|
|
| import cv2 |
| import numpy as np |
| import torch |
|
|
| sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src")) |
| sys.path.insert(0, os.path.dirname(__file__)) |
| from layered_belief import build_split_manifest, collect_decision_samples |
| from eval_l1_latent_understanding import FIELDS, balanced_limit, parse_offsets |
|
|
|
|
| def video_t_to_cell(video_t: float) -> int: |
| return max(0, int(round((float(video_t) - 0.125) / 0.25))) |
|
|
|
|
| def parse_cell_frames(text: str) -> List[int]: |
| vals = [int(x.strip()) for x in text.split(",") if x.strip()] |
| if not vals: |
| vals = [2] |
| bad = [v for v in vals if v < 0 or v > 3] |
| if bad: |
| raise ValueError(f"frame_in_cell must be in [0,3], got {bad}") |
| return vals |
|
|
|
|
| class RGBVideoIndex: |
| def __init__( |
| self, |
| processed_root: str, |
| history_offsets: Sequence[int], |
| frame_in_cell: Sequence[int], |
| height: int, |
| width: int, |
| backend: str = "auto", |
| max_open: int = 8, |
| ): |
| self.processed_root = processed_root |
| self.history_offsets = tuple(int(x) for x in history_offsets) |
| self.frame_in_cell = tuple(int(x) for x in frame_in_cell) |
| self.height = int(height) |
| self.width = int(width) |
| self.backend = self._resolve_backend(backend) |
| self.max_open = int(max_open) |
| self._caps: Dict[Tuple[str, int], Any] = {} |
| self._meta: Dict[Tuple[str, int], Dict[str, Any]] = {} |
|
|
| @staticmethod |
| def _resolve_backend(backend: str) -> str: |
| if backend != "auto": |
| return backend |
| try: |
| import decord |
| return "decord" |
| except Exception: |
| return "opencv" |
|
|
| def _path(self, boss: str, fight: int) -> str: |
| return os.path.join(self.processed_root, boss, f"video_fight{int(fight)}.mp4") |
|
|
| def _cap(self, boss: str, fight: int) -> Any | None: |
| key = (boss, int(fight)) |
| if key in self._caps: |
| return self._caps[key] |
| path = self._path(boss, fight) |
| if not os.path.exists(path): |
| return None |
| if self.backend == "decord": |
| try: |
| from decord import VideoReader, cpu |
| cap = VideoReader(path, ctx=cpu(0), width=self.width, height=self.height, num_threads=2) |
| except Exception: |
| return None |
| else: |
| cv2.setNumThreads(1) |
| cap = cv2.VideoCapture(path) |
| if not cap.isOpened(): |
| cap.release() |
| return None |
| if len(self._caps) >= self.max_open: |
| old_key = next(iter(self._caps)) |
| old = self._caps.pop(old_key) |
| if hasattr(old, "release"): |
| old.release() |
| self._caps[key] = cap |
| if self.backend == "decord": |
| self._meta[key] = { |
| "width": self.width, |
| "height": self.height, |
| "fps": float(cap.get_avg_fps()), |
| "frames": int(len(cap)), |
| "backend": "decord", |
| } |
| else: |
| self._meta[key] = { |
| "width": int(cap.get(cv2.CAP_PROP_FRAME_WIDTH)), |
| "height": int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT)), |
| "fps": float(cap.get(cv2.CAP_PROP_FPS)), |
| "frames": int(cap.get(cv2.CAP_PROP_FRAME_COUNT)), |
| "backend": "opencv", |
| } |
| return cap |
|
|
| def meta(self, boss: str, fight: int) -> Dict[str, Any] | None: |
| cap = self._cap(boss, fight) |
| if cap is None: |
| return None |
| return self._meta[(boss, int(fight))] |
|
|
| def close(self) -> None: |
| for cap in self._caps.values(): |
| if hasattr(cap, "release"): |
| cap.release() |
| self._caps.clear() |
|
|
| def read_frame(self, boss: str, fight: int, frame_idx: int) -> np.ndarray | None: |
| cap = self._cap(boss, fight) |
| if cap is None: |
| return None |
| meta = self._meta[(boss, int(fight))] |
| n_frames = int(meta["frames"]) |
| if n_frames <= 0: |
| return None |
| idx = min(max(0, int(frame_idx)), n_frames - 1) |
| if self.backend == "decord": |
| try: |
| return cap[idx].asnumpy() |
| except Exception: |
| return None |
| cap.set(cv2.CAP_PROP_POS_FRAMES, idx) |
| ok, bgr = cap.read() |
| if not ok or bgr is None: |
| return None |
| rgb = cv2.cvtColor(bgr, cv2.COLOR_BGR2RGB) |
| if rgb.shape[0] != self.height or rgb.shape[1] != self.width: |
| rgb = cv2.resize(rgb, (self.width, self.height), interpolation=cv2.INTER_AREA) |
| return rgb |
|
|
| def history(self, boss: str, fight: int, video_t: float) -> Tuple[torch.Tensor | None, List[int]]: |
| target_cell = video_t_to_cell(video_t) |
| frame_indices = [] |
| for off in self.history_offsets: |
| cell = max(0, target_cell + int(off)) |
| for in_cell in self.frame_in_cell: |
| frame_idx = 4 * cell + int(in_cell) |
| frame_indices.append(frame_idx) |
| if self.backend == "decord": |
| cap = self._cap(boss, fight) |
| if cap is None: |
| return None, frame_indices |
| meta = self._meta[(boss, int(fight))] |
| n_frames = int(meta["frames"]) |
| if n_frames <= 0: |
| return None, frame_indices |
| idxs = [min(max(0, int(i)), n_frames - 1) for i in frame_indices] |
| try: |
| batch = cap.get_batch(idxs).asnumpy() |
| except Exception: |
| return None, frame_indices |
| return torch.from_numpy(batch).permute(0, 3, 1, 2).contiguous().to(torch.uint8), frame_indices |
| frames = [] |
| for frame_idx in frame_indices: |
| rgb = self.read_frame(boss, fight, frame_idx) |
| if rgb is None: |
| return None, frame_indices |
| frames.append(torch.from_numpy(rgb).permute(2, 0, 1).contiguous()) |
| return torch.stack(frames, dim=0).to(torch.uint8), frame_indices |
|
|
|
|
| def label_distribution(labels: Dict[str, List[Any]]) -> Dict[str, Dict[str, int]]: |
| out = {} |
| for field, vals in labels.items(): |
| counts = Counter("__NONE__" if v is None else str(v) for v in vals) |
| out[field] = dict(sorted(counts.items(), key=lambda kv: kv[0])) |
| return out |
|
|
|
|
| def build_dataset( |
| samples: List[Dict[str, Any]], |
| index: RGBVideoIndex, |
| max_samples: int | None, |
| ) -> Tuple[torch.Tensor, Dict[str, List[Any]], List[Dict[str, Any]], Dict[str, Any]]: |
| kept = balanced_limit(samples, max_samples) |
| xs: List[torch.Tensor] = [] |
| labels = {field: [] for field in FIELDS} |
| used_samples = [] |
| missing = 0 |
| by_reason = Counter() |
| by_boss = Counter() |
| target_cells = [] |
| first_frame_indices = [] |
| video_meta_seen: Dict[str, Dict[str, Any]] = {} |
| for i, sample in enumerate(kept): |
| boss = sample["boss"] |
| fight = int(sample["fight"]) |
| meta = index.meta(boss, fight) |
| if meta is None: |
| missing += 1 |
| by_reason["missing_or_unreadable_video"] += 1 |
| continue |
| video_meta_seen[f"{boss}/fight{fight}"] = meta |
| hist, frame_indices = index.history(boss, fight, sample["belief"]["time"]) |
| if hist is None: |
| missing += 1 |
| by_reason["missing_frame"] += 1 |
| continue |
| xs.append(hist) |
| used_samples.append(sample) |
| by_boss[boss] += 1 |
| target_cells.append(video_t_to_cell(sample["belief"]["time"])) |
| first_frame_indices.append(frame_indices[0] if frame_indices else None) |
| for field in FIELDS: |
| labels[field].append(sample["belief"][field]) |
| if (i + 1) % 1000 == 0: |
| print(f"processed {i + 1}/{len(kept)} kept={len(xs)} missing={missing}", flush=True) |
| index.close() |
| if not xs: |
| raise RuntimeError("no RGB samples were loaded") |
| X = torch.stack(xs, dim=0).contiguous() |
| audit = { |
| "requested": len(kept), |
| "kept": len(xs), |
| "missing": int(missing), |
| "missing_by_reason": dict(by_reason), |
| "kept_by_boss": dict(by_boss), |
| "label_distribution": label_distribution(labels), |
| "target_cell_min": int(min(target_cells)) if target_cells else None, |
| "target_cell_max": int(max(target_cells)) if target_cells else None, |
| "first_history_frame_min": int(min(x for x in first_frame_indices if x is not None)) if first_frame_indices else None, |
| "first_history_frame_max": int(max(x for x in first_frame_indices if x is not None)) if first_frame_indices else None, |
| "video_meta_examples": dict(list(video_meta_seen.items())[:6]), |
| } |
| return X, labels, used_samples, audit |
|
|
|
|
| _WORKER: Dict[str, Any] = {} |
|
|
|
|
| def _worker_init(processed_root, offsets, frame_in_cell, height, width, backend, max_open) -> None: |
| _WORKER["index"] = RGBVideoIndex( |
| processed_root, offsets, frame_in_cell, height, width, backend, max_open |
| ) |
|
|
|
|
| def _worker_fight(task): |
| """Decode every sample that lives in one fight, so its video opens once.""" |
| boss, fight, items = task |
| index = _WORKER["index"] |
| meta = index.meta(boss, fight) |
| if meta is None: |
| return boss, fight, None, [(pos, "missing_or_unreadable_video") for pos, _ in items], None |
| out, failed = [], [] |
| for pos, video_t in items: |
| hist, frame_indices = index.history(boss, fight, video_t) |
| if hist is None: |
| failed.append((pos, "missing_frame")) |
| continue |
| out.append((pos, hist.numpy(), frame_indices)) |
| return boss, fight, out, failed, meta |
|
|
|
|
| def build_dataset_parallel( |
| samples: List[Dict[str, Any]], |
| max_samples: int | None, |
| workers: int, |
| processed_root: str, |
| offsets, |
| frame_in_cell, |
| height: int, |
| width: int, |
| backend: str, |
| max_open: int, |
| ) -> Tuple[torch.Tensor, Dict[str, List[Any]], List[Dict[str, Any]], Dict[str, Any]]: |
| """Same result as build_dataset, but fights are decoded across processes. |
| |
| Decoding is the whole cost here and each fight is an independent video, so |
| the serial path leaves 127 of this box's 128 cores idle. Results are keyed |
| back to their original position, so output order matches the serial path. |
| """ |
| kept = balanced_limit(samples, max_samples) |
| by_fight: Dict[Tuple[str, int], List[Tuple[int, float]]] = defaultdict(list) |
| for pos, sample in enumerate(kept): |
| by_fight[(sample["boss"], int(sample["fight"]))].append( |
| (pos, sample["belief"]["time"]) |
| ) |
| tasks = [(boss, fight, items) for (boss, fight), items in by_fight.items()] |
| print(f"parallel: {len(kept)} samples across {len(tasks)} fights, {workers} workers", flush=True) |
|
|
| frames: Dict[int, Any] = {} |
| first_frame_by_pos: Dict[int, Any] = {} |
| by_reason: Counter = Counter() |
| video_meta_seen: Dict[str, Dict[str, Any]] = {} |
|
|
| ctx = mp.get_context("spawn") |
| with ProcessPoolExecutor( |
| max_workers=workers, |
| mp_context=ctx, |
| initializer=_worker_init, |
| initargs=(processed_root, offsets, frame_in_cell, height, width, backend, max_open), |
| ) as pool: |
| futures = [pool.submit(_worker_fight, t) for t in tasks] |
| for done, fut in enumerate(as_completed(futures), 1): |
| boss, fight, out, failed, meta = fut.result() |
| if meta is not None: |
| video_meta_seen[f"{boss}/fight{fight}"] = meta |
| for pos, reason in failed: |
| by_reason[reason] += 1 |
| for pos, arr, frame_indices in (out or []): |
| frames[pos] = arr |
| first_frame_by_pos[pos] = frame_indices[0] if frame_indices else None |
| if done % 50 == 0 or done == len(tasks): |
| print( |
| f"fights {done}/{len(tasks)} kept={len(frames)} missing={sum(by_reason.values())}", |
| flush=True, |
| ) |
|
|
| order = sorted(frames) |
| if not order: |
| raise RuntimeError("no RGB samples were loaded") |
|
|
| X = torch.from_numpy(np.stack([frames[pos] for pos in order], axis=0)).contiguous() |
| labels = {field: [kept[pos]["belief"][field] for pos in order] for field in FIELDS} |
| used_samples = [kept[pos] for pos in order] |
| by_boss = Counter(s["boss"] for s in used_samples) |
| target_cells = [video_t_to_cell(s["belief"]["time"]) for s in used_samples] |
| first_frame_indices = [first_frame_by_pos[pos] for pos in order] |
| missing = sum(by_reason.values()) |
| audit = { |
| "requested": len(kept), |
| "kept": len(order), |
| "missing": int(missing), |
| "missing_by_reason": dict(by_reason), |
| "kept_by_boss": dict(by_boss), |
| "label_distribution": label_distribution(labels), |
| "target_cell_min": int(min(target_cells)) if target_cells else None, |
| "target_cell_max": int(max(target_cells)) if target_cells else None, |
| "first_history_frame_min": int(min(x for x in first_frame_indices if x is not None)) if first_frame_indices else None, |
| "first_history_frame_max": int(max(x for x in first_frame_indices if x is not None)) if first_frame_indices else None, |
| "video_meta_examples": dict(list(video_meta_seen.items())[:6]), |
| "workers": workers, |
| } |
| return X, labels, used_samples, audit |
|
|
|
|
| def main() -> None: |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--manifest", default="out/layered/manifest_latent_covered.json") |
| ap.add_argument("--processed_root", default="data/processed") |
| ap.add_argument("--split", choices=["train", "val", "test"], required=True) |
| ap.add_argument("--boss", default=None, help="Optional entity/boss id to filter; defaults to all manifest entries.") |
| ap.add_argument("--history_offsets", default="-8,-4,-2,-1") |
| ap.add_argument("--frame_in_cell", default="2") |
| ap.add_argument("--height", type=int, default=192) |
| ap.add_argument("--width", type=int, default=336) |
| ap.add_argument("--backend", choices=["auto", "decord", "opencv"], default="auto") |
| ap.add_argument("--max_samples", type=int, default=None) |
| ap.add_argument("--max_open", type=int, default=8) |
| ap.add_argument("--allow_current_frame", action="store_true") |
| ap.add_argument( |
| "--workers", |
| type=int, |
| default=1, |
| help="Decode fights across this many processes. 1 keeps the original serial path.", |
| ) |
| ap.add_argument("--out", required=True) |
| args = ap.parse_args() |
|
|
| offsets = parse_offsets(args.history_offsets) |
| if any(int(x) >= 0 for x in offsets) and not args.allow_current_frame: |
| raise ValueError("raw-RGB causal cache requires strictly negative history offsets") |
| if any(int(x) > 0 for x in offsets): |
| raise ValueError("raw-RGB cache never allows positive future offsets") |
| frame_in_cell = parse_cell_frames(args.frame_in_cell) |
|
|
| if os.path.exists(args.manifest): |
| with open(args.manifest, encoding="utf-8") as f: |
| manifest: Dict[str, Any] = json.load(f) |
| else: |
| manifest = build_split_manifest() |
| samples = collect_decision_samples(manifest, args.split, args.boss) |
| if args.workers > 1: |
| probe = RGBVideoIndex( |
| args.processed_root, offsets, frame_in_cell, |
| args.height, args.width, args.backend, args.max_open, |
| ) |
| resolved_backend = probe.backend |
| probe.close() |
| X, labels, used_samples, audit = build_dataset_parallel( |
| samples, |
| args.max_samples, |
| args.workers, |
| args.processed_root, |
| offsets, |
| frame_in_cell, |
| args.height, |
| args.width, |
| args.backend, |
| args.max_open, |
| ) |
| else: |
| index = RGBVideoIndex( |
| args.processed_root, |
| offsets, |
| frame_in_cell, |
| args.height, |
| args.width, |
| args.backend, |
| args.max_open, |
| ) |
| X, labels, used_samples, audit = build_dataset(samples, index, args.max_samples) |
| resolved_backend = index.backend |
| payload = { |
| "rgb": X, |
| "labels": labels, |
| "samples": used_samples, |
| "history_offsets": list(offsets), |
| "frame_in_cell": list(frame_in_cell), |
| "rgb_height": args.height, |
| "rgb_width": args.width, |
| "rgb_source": "raw_processed_fight_video", |
| "video_backend": resolved_backend, |
| "manifest": args.manifest, |
| "processed_root": args.processed_root, |
| "split": args.split, |
| "boss": args.boss or "all", |
| "max_samples": args.max_samples, |
| "allow_current_frame": args.allow_current_frame, |
| "n_samples": int(len(X)), |
| "audit": audit, |
| } |
| os.makedirs(os.path.dirname(args.out), exist_ok=True) |
| torch.save(payload, args.out) |
| summary = { |
| "out": args.out, |
| "split": args.split, |
| "boss": args.boss or "all", |
| "rgb_source": payload["rgb_source"], |
| "video_backend": payload["video_backend"], |
| "history_offsets": list(offsets), |
| "frame_in_cell": list(frame_in_cell), |
| "allow_current_frame": args.allow_current_frame, |
| "shape": list(X.shape), |
| "dtype": str(X.dtype), |
| "audit": audit, |
| } |
| print(json.dumps(summary, ensure_ascii=False, indent=2), flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| main() |
|
|