File size: 15,292 Bytes
242cc21 | 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 | """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 # ffmpeg-backed decode (works in this venv; avoids the torchcodec dependency)
from .actions import KeyVocab
from .batch import VideoActionBatch
from .clips import compute_stride
from .training_loader import ClipMeta, _collate # reuse MIRA's collate + meta type
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") # (H, W, 3) uint8
if i >= hi:
break
arr = np.stack([grabbed[i] for i in frame_indices]) # (T, H, W, 3)
t = torch.from_numpy(arr).permute(0, 3, 1, 2).contiguous() # (T, C, H, W) uint8
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 held-key characters (per-frame `actions` string), stable multi-hot order. '-' == no input.
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()
# --- ranged tar-member fetch for the FULL (tarred) split ------------------------------------------
# A round's POVs are scattered across many ~2GB tars. We open each tar over HfFileSystem (a seekable,
# HTTP-range-backed file) with tarfile mode="r", which walks the member headers via small ranged
# reads ONCE per shard, then extracts a single member with one ranged read of just its bytes — so we
# never download a whole 2GB tar. Open TarFiles are cached per shard.
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") # seekable, range-backed
tf = tarfile.open(fileobj=f, mode="r") # header walk (ranged) once
members = {m.name.rsplit("/", 1)[-1]: m for m in tf.getmembers()}
if len(_TAR_CACHE) > 24: # bound open handles
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)
# group rows into (match_id, round_number) -> [row, ...] sorted by player_index
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)
# keep only groups with enough perspectives to form an n_players block
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: # full split: ranged tar members
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: # sample split: individual files
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) # align at frame 0, truncate to shortest
n_steps = n_src // self.stride
if n_steps < self.clip_len:
return
# per-perspective full-round action arrays, then slice per clip window
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)
# Decode every target-fps frame of the clip ONCE per perspective, then slice windows from it.
# (Decoding per window re-decoded from the start of the mp4 each time -> O(n^2) and starved
# the GPU; this is a single linear pass.)
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] # (T,C,H,W) uint8, already decoded
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: # a bad clip shouldn't kill the epoch
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) # (1, T, n_keys)
at.mouse_movements = mouse.unsqueeze(0).to(torch.float32) # (1, T, 2)
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)
# sample/ paths are relative to sample/; make them repo-relative
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"), # full split = ranged tar-member fetch; sample = individual files
)
return DataLoader(
ds,
batch_size=batch_size * n_players, # n_players contiguous rows per sample-block
num_workers=num_workers,
collate_fn=_collate,
drop_last=True,
prefetch_factor=prefetch_factor if num_workers else None,
)
|