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,
    )