| """Streaming CS2-10k dataset for MIRA (no local dataset build). |
| |
| RekaAI/CS2-10k is a WebDataset of first-person CS2 rounds. `index.parquet` lists every clip with: |
| match_id, video_path, parquet_path, map, player_index, round_number, team, fps(48), total_time, |
| width, height, match_num_clips, match_num_rounds, match_num_players, shard |
| Grouping key for MIRA's multi-perspective samples is (match_id, round_number): fixing both and |
| varying player_index gives up to `match_num_players` (=10) synchronized first-person views of the |
| same round — the CS2 analogue of MIRA's 4 Rocket League views. |
| |
| This module streams those groups straight from the Hub and yields MIRA's own |
| `(VideoActionBatch, list[ClipMeta])` batches, so it drops into MIRA's trainer in place of |
| `create_loader` with zero on-disk dataset. It reuses MIRA's `decode_frames`, `KeyVocab`, |
| `ActionTensors`, `ClipMeta`, `VideoActionBatch`, and `_collate`. |
| |
| Two important, honest caveats: |
| * ALIGNMENT: CS2 gives no shared per-frame tick. Each round's per-player clips are 48fps and |
| (empirically) start at round start, so we align at frame 0 and truncate a group to the shortest |
| clip. Refine later with position/rotation cross-checks if needed. |
| * SCALE / STREAMING UNIT: the 63TB main split is tar-sharded and a round's POVs are scattered |
| across many shards, so complete groups are NOT co-located in one tar. This loader therefore |
| fetches clips as INDIVIDUAL files by path (works today for the untarred `sample/` split — 3 full |
| matches — and for any individual-file mirror). Training on the full tarred split at scale needs |
| either re-sharding by (match,round) or a per-group multi-shard fetch; that's a separate step. |
| """ |
|
|
| from __future__ import annotations |
|
|
| import io |
| import random |
| from collections import defaultdict |
| from itertools import count |
| from typing import Any, Iterator |
|
|
| import numpy as np |
| import torch |
| from torch.utils.data import DataLoader, IterableDataset, get_worker_info |
|
|
| from huggingface_hub import hf_hub_download |
| import pandas as pd |
|
|
| import av |
|
|
| from .actions import KeyVocab |
| from .batch import VideoActionBatch |
| from .clips import compute_stride |
| from .training_loader import ClipMeta, _collate |
|
|
|
|
| def _decode_av(mp4_bytes: bytes, frame_indices: list[int], frame_size: tuple[int, int] | None): |
| """Decode the given source-frame indices from mp4 bytes to (T, C, H, W) uint8, via PyAV.""" |
| want = set(frame_indices) |
| hi = max(frame_indices) |
| grabbed: dict[int, np.ndarray] = {} |
| with av.open(io.BytesIO(mp4_bytes)) as container: |
| for i, frame in enumerate(container.decode(container.streams.video[0])): |
| if i in want: |
| grabbed[i] = frame.to_ndarray(format="rgb24") |
| if i >= hi: |
| break |
| arr = np.stack([grabbed[i] for i in frame_indices]) |
| t = torch.from_numpy(arr).permute(0, 3, 1, 2).contiguous() |
| if frame_size is not None: |
| t = torch.nn.functional.interpolate( |
| t.float(), size=frame_size, mode="bilinear", align_corners=False |
| ).clamp(0, 255).to(torch.uint8) |
| return t |
|
|
| |
| CS2_KEYS: tuple[str, ...] = ("W", "A", "S", "D", "J", "C", "R", "V", "[", "]") |
|
|
| REPO_ID = "RekaAI/CS2-10k" |
|
|
|
|
| def _rank_world() -> tuple[int, int]: |
| if torch.distributed.is_available() and torch.distributed.is_initialized(): |
| return torch.distributed.get_rank(), torch.distributed.get_world_size() |
| return 0, 1 |
|
|
|
|
| def _fetch_bytes(path: str) -> bytes: |
| """Fetch an INDIVIDUAL file by repo path (works for the untarred `sample/` split).""" |
| local = hf_hub_download(REPO_ID, path, repo_type="dataset") |
| with open(local, "rb") as f: |
| return f.read() |
|
|
|
|
| |
| |
| |
| |
| |
| import tarfile |
|
|
| _HF_FS = None |
| _TAR_CACHE: dict[str, tuple] = {} |
|
|
|
|
| def _hf_fs(): |
| global _HF_FS |
| if _HF_FS is None: |
| from huggingface_hub import HfFileSystem |
| _HF_FS = HfFileSystem() |
| return _HF_FS |
|
|
|
|
| def _open_shard(shard: str): |
| if shard not in _TAR_CACHE: |
| f = _hf_fs().open(f"datasets/{REPO_ID}/{shard}", "rb") |
| tf = tarfile.open(fileobj=f, mode="r") |
| members = {m.name.rsplit("/", 1)[-1]: m for m in tf.getmembers()} |
| if len(_TAR_CACHE) > 24: |
| old = next(iter(_TAR_CACHE)) |
| try: |
| _TAR_CACHE.pop(old)[0].close() |
| except Exception: |
| pass |
| _TAR_CACHE[shard] = (tf, members) |
| return _TAR_CACHE[shard] |
|
|
|
|
| def _fetch_from_tar(shard: str, repo_path: str) -> bytes: |
| tf, members = _open_shard(shard) |
| m = members[repo_path.rsplit("/", 1)[-1]] |
| return tf.extractfile(m).read() |
|
|
|
|
| def _build_action_arrays( |
| frame_data: np.ndarray, vocab: KeyVocab, stride: int |
| ) -> tuple[torch.Tensor, torch.Tensor]: |
| """From a clip's 48fps per-frame records build downsampled (n_steps, n_keys) int32 multi-hot key |
| presses (OR-ed over each stride window) and (n_steps, 2) float32 mean mouse deltas.""" |
| n_keys = len(vocab) |
| n_steps = len(frame_data) // stride |
| keys = torch.zeros((n_steps, n_keys), dtype=torch.int32) |
| mouse = torch.zeros((n_steps, 2), dtype=torch.float32) |
| for s in range(n_steps): |
| mx = my = 0.0 |
| for j in range(s * stride, (s + 1) * stride): |
| rec = frame_data[j] |
| a = rec["actions"] |
| if a and a != "-": |
| for ch in a: |
| idx = vocab._index.get(ch) |
| if idx is not None: |
| keys[s, idx] = 1 |
| mx += float(rec["mouse_x_delta"]) |
| my += float(rec["mouse_y_delta"]) |
| mouse[s, 0] = mx / stride |
| mouse[s, 1] = my / stride |
| return keys, mouse |
|
|
|
|
| class Cs2StreamingDataset(IterableDataset): |
| """Streams (match_id, round_number) POV groups from CS2-10k as MIRA per-perspective samples. |
| |
| Yields dicts shaped exactly like MIRA's training_loader `_decode_sample` output |
| ({"video","actions","metadata"}), so `_collate` turns a batch into `(VideoActionBatch, [ClipMeta])`. |
| """ |
|
|
| def __init__( |
| self, |
| index: pd.DataFrame, |
| *, |
| vocab: KeyVocab, |
| clip_len: int = 16, |
| target_fps: int = 16, |
| source_fps: int = 48, |
| n_players: int = 4, |
| frame_size: tuple[int, int] | None = None, |
| mouse_sensitivity: float = 1.0, |
| shuffle: bool = True, |
| infinite: bool = True, |
| seed: int = 2025, |
| from_tar: bool = False, |
| ) -> None: |
| super().__init__() |
| self.index = index |
| self.from_tar = from_tar |
| self.vocab = vocab |
| self.clip_len = clip_len |
| self.target_fps = target_fps |
| self.source_fps = source_fps |
| self.n_players = n_players |
| self.frame_size = frame_size |
| self.mouse_sensitivity = mouse_sensitivity |
| self.shuffle = shuffle |
| self.infinite = infinite |
| self.seed = seed |
| self.stride = compute_stride(source_fps, target_fps) |
| |
| groups: dict[tuple[str, int], list[dict]] = defaultdict(list) |
| for r in index.to_dict("records"): |
| groups[(r["match_id"], int(r["round_number"]))].append(r) |
| |
| self.groups: list[list[dict]] = [ |
| sorted(v, key=lambda r: r["player_index"]) |
| for v in groups.values() |
| if len(v) >= n_players |
| ] |
|
|
| def _my_groups(self) -> list[list[dict]]: |
| rank, world = _rank_world() |
| info = get_worker_info() |
| wid, nw = (info.id, info.num_workers) if info else (0, 1) |
| return self.groups[rank::world][wid::nw] |
|
|
| def _load_group(self, rows: list[dict]) -> tuple[list[bytes], list[np.ndarray], list[dict]]: |
| import io as _io |
| vids, fdata, meta = [], [], [] |
| for r in rows[: self.n_players]: |
| if self.from_tar: |
| vids.append(_fetch_from_tar(r["shard"], r["video_path"])) |
| pq = pd.read_parquet(_io.BytesIO(_fetch_from_tar(r["shard"], r["parquet_path"]))) |
| else: |
| vids.append(_fetch_bytes(r["video_path"])) |
| pq = pd.read_parquet(hf_hub_download(REPO_ID, r["parquet_path"], repo_type="dataset")) |
| fdata.append(pq.iloc[0]["frame_data"]) |
| meta.append(r) |
| return vids, fdata, meta |
|
|
| def _samples_from_group(self, rows: list[dict]) -> Iterator[dict[str, Any]]: |
| vids, fdata, meta = self._load_group(rows) |
| n_src = min(len(fd) for fd in fdata) |
| n_steps = n_src // self.stride |
| if n_steps < self.clip_len: |
| return |
| |
| key_arrs, mouse_arrs = [], [] |
| for fd in fdata: |
| k, m = _build_action_arrays(fd[: n_steps * self.stride], self.vocab, self.stride) |
| key_arrs.append(k) |
| mouse_arrs.append(m) |
| |
| |
| |
| all_idx = [s * self.stride for s in range(n_steps)] |
| full = [_decode_av(vids[p], all_idx, self.frame_size) for p in range(self.n_players)] |
| clip_id = 0 |
| for start in range(0, n_steps - self.clip_len + 1, self.clip_len): |
| frame_indices = all_idx[start : start + self.clip_len] |
| for p in range(self.n_players): |
| video = full[p][start : start + self.clip_len] |
| act = _make_action_tensors( |
| key_arrs[p][start : start + self.clip_len], |
| mouse_arrs[p][start : start + self.clip_len], |
| self.vocab, |
| self.mouse_sensitivity, |
| ) |
| yield { |
| "video": video, |
| "actions": act, |
| "metadata": ClipMeta( |
| match_id=meta[p]["match_id"], |
| perspective=p, |
| player_id=int(meta[p]["player_index"]), |
| clip_id=clip_id, |
| chunk_idx=int(meta[p]["round_number"]), |
| frame_indices=list(frame_indices), |
| ), |
| } |
| clip_id += 1 |
|
|
| def __iter__(self) -> Iterator[dict[str, Any]]: |
| rank, _ = _rank_world() |
| info = get_worker_info() |
| rng = random.Random(self.seed + rank * 1024 + (info.id if info else 0)) |
| for _epoch in count() if self.infinite else range(1): |
| order = self._my_groups() |
| if self.shuffle: |
| rng.shuffle(order) |
| for rows in order: |
| try: |
| yield from self._samples_from_group(rows) |
| except Exception as e: |
| print(f"[cs2_stream] skipping group {rows[0]['match_id']}: {e}") |
|
|
|
|
| def _make_action_tensors(keys: torch.Tensor, mouse: torch.Tensor, vocab: KeyVocab, sens: float): |
| """Build a MIRA ActionTensors (batch=1) from a clip window's key/mouse tensors.""" |
| from mira.world_model.actions_config import ActionConfig, ActionTensors |
|
|
| cfg = ActionConfig(valid_keys=list(vocab.keys)) |
| at = ActionTensors(config=cfg, batch_size=1) |
| at.key_presses = keys.unsqueeze(0).to(torch.int32) |
| at.mouse_movements = mouse.unsqueeze(0).to(torch.float32) |
| at.game_mouse_sensitivity = torch.full((1,), float(sens), dtype=torch.float32) |
| return at |
|
|
|
|
| def load_cs2_index(subset: str = "sample") -> pd.DataFrame: |
| """Return the CS2 index as a DataFrame. `subset='sample'` uses the untarred 3-match `sample/` |
| split whose clips are individually fetchable (recommended for streaming today); `subset='full'` |
| uses the root index (tarred; see the module docstring caveat before using at scale).""" |
| if subset == "sample": |
| p = hf_hub_download(REPO_ID, "sample/index.parquet", repo_type="dataset") |
| df = pd.read_parquet(p) |
| |
| for col in ("video_path", "parquet_path"): |
| if col in df.columns: |
| df[col] = df[col].apply(lambda x: x if x.startswith("sample/") else f"sample/{x}") |
| return df |
| p = hf_hub_download(REPO_ID, "index.parquet", repo_type="dataset") |
| return pd.read_parquet(p) |
|
|
|
|
| def create_cs2_loader( |
| *, |
| subset: str = "sample", |
| index: pd.DataFrame | None = None, |
| maps: list[str] | None = None, |
| clip_len: int = 16, |
| target_fps: int = 16, |
| n_players: int = 4, |
| batch_size: int = 4, |
| num_workers: int = 0, |
| frame_size: tuple[int, int] | None = None, |
| valid_keys: list[str] | None = None, |
| shuffle: bool = True, |
| infinite: bool = True, |
| seed: int = 2025, |
| prefetch_factor: int = 2, |
| ) -> DataLoader: |
| """Drop-in replacement for MIRA's `create_loader`, streaming CS2-10k from the Hub (no disk). |
| |
| Yields `(VideoActionBatch, list[ClipMeta])`. `n_players` perspectives per (match, round) are |
| grouped contiguously so the collate stacks them into one multi-perspective block, exactly as |
| MIRA's own loader does. |
| """ |
| if index is None: |
| index = load_cs2_index(subset) |
| if maps: |
| index = index[index["map"].isin(maps)] |
| vocab = KeyVocab(tuple(valid_keys) if valid_keys else CS2_KEYS, on_unknown="ignore") |
| ds = Cs2StreamingDataset( |
| index, |
| vocab=vocab, |
| clip_len=clip_len, |
| target_fps=target_fps, |
| n_players=n_players, |
| frame_size=frame_size, |
| shuffle=shuffle, |
| infinite=infinite, |
| seed=seed, |
| from_tar=(subset == "full"), |
| ) |
| return DataLoader( |
| ds, |
| batch_size=batch_size * n_players, |
| num_workers=num_workers, |
| collate_fn=_collate, |
| drop_last=True, |
| prefetch_factor=prefetch_factor if num_workers else None, |
| ) |
|
|