from __future__ import annotations import gc import threading from typing import Any import numpy as np import torch from torch.utils.data import Dataset try: from src.data_pipeline.fast_dataset import FastDataset except ImportError: from ..data_pipeline.fast_dataset import FastDataset # type: ignore def _numpy_dtype(name: str) -> np.dtype: if str(name).lower() in {"float16", "fp16", "half"}: return np.dtype(np.float16) if str(name).lower() in {"float32", "fp32", "single"}: return np.dtype(np.float32) return np.dtype(name) class FastRAMChunkedDataset(Dataset): """ RAM chunk wrapper for FastDataset. It keeps the public sample format identical to FastDataset, but each chunk caches unique .dat rows once and reconstructs time windows by local indices. """ def __init__( self, base: FastDataset, chunk_ram_gb: float = 40.0, cache_dtype: str | np.dtype = "float16", read_block_rows: int = 64, async_prefetch: bool = True, chunk_order: str = "sequential", verbose: bool = True, ): if not isinstance(base, FastDataset): raise TypeError("FastRAMChunkedDataset only supports FastDataset") self.base = base self.chunk_ram_gb = float(chunk_ram_gb) self.cache_dtype = _numpy_dtype(str(cache_dtype)) self.read_block_rows = max(1, int(read_block_rows)) self.async_prefetch = bool(async_prefetch) self.chunk_order = str(chunk_order) self.verbose = bool(verbose) self.target_bytes = int(self.chunk_ram_gb * (1024 ** 3)) self.chunks = self._build_chunks() self.current: dict[str, Any] | None = None self.current_chunk_id: int | None = None self._preload_thread: threading.Thread | None = None self._preload_result: dict[str, Any] | None = None self._preload_error: BaseException | None = None self._preload_chunk_id: int | None = None self._preload_done = threading.Event() self._preload_lock = threading.Lock() if self.verbose: print( f"[FastRAMChunkedDataset] samples={len(self.base)}, " f"target={self.chunk_ram_gb:.2f} GB, chunks={len(self.chunks)}, " f"chunk_order={self.chunk_order}" ) for i, chunk in enumerate(self.chunks[:10]): print(f" chunk {i:03d}: {self._chunk_desc(chunk)}") if len(self.chunks) > 10: print(f" ... {len(self.chunks) - 10} more chunks") @property def num_chunks(self) -> int: return len(self.chunks) def _source_cache_dtype(self, source: str) -> np.dtype: store_dtype = self.base.stores[source].dtype if np.issubdtype(store_dtype, np.floating): return self.cache_dtype return store_dtype def _row_bytes(self, source: str) -> int: store = self.base.stores[source] return int(np.prod(store.row_shape) * self._source_cache_dtype(source).itemsize) def _sample_refs(self, sample_time: str) -> dict[str, set[int]]: refs: dict[str, set[int]] = {} for name in self.base.required_inputs: source = self.base._catalog_source_name(name) refs.setdefault(source, set()) for ts in self.base._input_times(sample_time): refs[source].add(self.base._idx_for(ts, source)) for label in self.base.required_labels: source = self.base._catalog_source_name(label) refs.setdefault(source, set()) if label == "bt": if self.base.use_bt_mask: mask_source = self.base.bt_mask_source refs.setdefault(mask_source, set()).add(self.base._idx_for(sample_time, mask_source)) for ts in self.base._bt_times(sample_time): refs[source].add(self.base._idx_for(ts, source)) else: refs[source].add(self.base._idx_for(sample_time, source)) return refs def _chunk_desc(self, chunk: dict[str, Any]) -> str: start = int(chunk["start"]) end = int(chunk["end"]) n = int(len(chunk["indices"])) est_gb = float(chunk["est_gb"]) if chunk.get("wrapped"): wrap_count = int(chunk.get("wrap_count", 0)) return f"samples [{start}:{len(self.base.samples)}) + [0:{wrap_count}) n={n} est={est_gb:.2f} GB" return f"samples [{start}:{end}) n={n} est={est_gb:.2f} GB" def _refs_for_indices(self, sample_indices: list[int] | np.ndarray) -> dict[str, set[int]]: refs: dict[str, set[int]] = {} for sample_idx in sample_indices: ts = self.base.samples[int(sample_idx)] for source, indices in self._sample_refs(ts).items(): refs.setdefault(source, set()).update(indices) return refs def _make_chunk(self, sample_indices: list[int] | np.ndarray, start: int, end: int, wrapped: bool = False) -> dict[str, Any]: indices = np.asarray(sample_indices, dtype=np.int64) refs = self._refs_for_indices(indices) return { "start": int(start), "end": int(end), "indices": indices, "est_gb": self._estimate_bytes(refs, len(indices)) / (1024 ** 3), "wrapped": bool(wrapped), "wrap_count": int(np.sum(indices < start)) if wrapped else 0, } def _build_chunks(self) -> list[dict[str, Any]]: chunks: list[dict[str, Any]] = [] start = 0 refs: dict[str, set[int]] = {} for i, sample_time in enumerate(self.base.samples): sample_refs = self._sample_refs(sample_time) for source, indices in sample_refs.items(): refs.setdefault(source, set()).update(indices) n_samples = i - start + 1 est_bytes = self._estimate_bytes(refs, n_samples) if est_bytes > self.target_bytes and i > start: end = i chunks.append(self._make_chunk(list(range(start, end)), start=start, end=end)) start = i refs = sample_refs if start < len(self.base.samples): chunks.append(self._make_chunk(list(range(start, len(self.base.samples))), start=start, end=len(self.base.samples))) if self.chunk_order == "circular": self._fill_last_chunk_circular(chunks) return chunks def _fill_last_chunk_circular(self, chunks: list[dict[str, Any]]) -> None: if len(chunks) <= 1 or len(self.base.samples) == 0: return last = chunks[-1] if last.get("wrapped"): return previous_sizes = [len(chunk["indices"]) for chunk in chunks[:-1]] target_samples = int(round(float(np.median(previous_sizes)))) if previous_sizes else len(last["indices"]) missing = max(0, target_samples - len(last["indices"])) if missing <= 0: return fill_count = min(missing, len(self.base.samples) - len(last["indices"])) if fill_count <= 0: return indices = np.concatenate([last["indices"], np.arange(fill_count, dtype=np.int64)]) chunks[-1] = self._make_chunk(indices, start=int(last["start"]), end=len(self.base.samples), wrapped=True) def _estimate_bytes(self, refs: dict[str, set[int]], n_samples: int) -> int: rows = sum(len(indices) * self._row_bytes(source) for source, indices in refs.items()) per_sample_refs = n_samples * 16 * ( len(self.base.required_inputs) * len(self.base.input_offsets) + len(self.base.required_labels) * max(1, len(self.base.bt_offsets)) ) return int(rows + per_sample_refs) def _load_source_rows(self, source: str, indices: np.ndarray) -> np.ndarray: store = self.base.stores[source] dtype = self._source_cache_dtype(source) out = np.empty((len(indices), *store.row_shape), dtype=dtype) for start in range(0, len(indices), self.read_block_rows): end = min(start + self.read_block_rows, len(indices)) out[start:end] = store._mm[indices[start:end]].astype(dtype, copy=False) return out def _load_chunk(self, chunk_id: int) -> dict[str, Any]: chunk_id = int(chunk_id) % len(self.chunks) chunk = self.chunks[chunk_id] sample_indices = chunk["indices"] sample_times = [self.base.samples[int(sample_idx)] for sample_idx in sample_indices] refs: dict[str, set[int]] = {} for ts in sample_times: for source, indices in self._sample_refs(ts).items(): refs.setdefault(source, set()).update(indices) if self.verbose: print( f"\n[FastRAMChunkedDataset] loading chunk {chunk_id}/{len(self.chunks)-1} " f"| {self._chunk_desc(chunk)}" ) arrays: dict[str, np.ndarray] = {} local_maps: dict[str, dict[int, int]] = {} for source, idx_set in sorted(refs.items()): indices = np.array(sorted(idx_set), dtype=np.int64) arrays[source] = self._load_source_rows(source, indices) local_maps[source] = {int(idx): i for i, idx in enumerate(indices)} real_gb = sum(arr.nbytes for arr in arrays.values()) / (1024 ** 3) if self.verbose: detail = ", ".join(f"{src}={arr.shape}{arr.dtype}" for src, arr in arrays.items()) print(f"[FastRAMChunkedDataset] loaded chunk {chunk_id} | {detail} | real={real_gb:.2f} GB") return { "chunk_id": chunk_id, "sample_start": int(chunk["start"]), "sample_end": int(chunk["end"]), "sample_indices": sample_indices, "sample_times": sample_times, "arrays": arrays, "local_maps": local_maps, "real_gb": real_gb, } def load_chunk_sync(self, chunk_id: int, free_current_before_load: bool = False) -> None: if free_current_before_load: self.current = None gc.collect() self.current = self._load_chunk(chunk_id) self.current_chunk_id = int(self.current["chunk_id"]) gc.collect() def start_preload(self, chunk_id: int) -> bool: if not self.async_prefetch or len(self.chunks) <= 1: return False chunk_id = int(chunk_id) % len(self.chunks) with self._preload_lock: if self._preload_thread is not None and self._preload_thread.is_alive(): return False if self._preload_result is not None: return False self._preload_result = None self._preload_error = None self._preload_chunk_id = chunk_id self._preload_done.clear() def _worker() -> None: try: result = self._load_chunk(chunk_id) with self._preload_lock: self._preload_result = result self._preload_error = None except BaseException as exc: with self._preload_lock: self._preload_result = None self._preload_error = exc finally: self._preload_done.set() self._preload_thread = threading.Thread(target=_worker, daemon=True) self._preload_thread.start() if self.verbose: print(f"[FastRAMChunkedDataset] background preload started: chunk {chunk_id}") return True def swap_if_preload_ready(self) -> bool: if not self._preload_done.is_set(): if self.verbose and self._preload_chunk_id is not None: print(f"[FastRAMChunkedDataset] preload not ready yet: chunk {self._preload_chunk_id}") return False with self._preload_lock: if self._preload_error is not None: raise RuntimeError(f"background preload failed: {self._preload_error}") from self._preload_error if self._preload_result is None: return False old = self.current old_id = self.current_chunk_id self.current = self._preload_result self.current_chunk_id = int(self.current["chunk_id"]) self._preload_result = None self._preload_error = None self._preload_chunk_id = None self._preload_done.clear() if self._preload_thread is not None: self._preload_thread.join(timeout=0) self._preload_thread = None del old gc.collect() if self.verbose: print(f"[FastRAMChunkedDataset] swapped chunk {old_id} -> {self.current_chunk_id}") return True def wait_for_preload_and_swap(self) -> bool: thread = self._preload_thread if thread is not None and thread.is_alive(): thread.join() return self.swap_if_preload_ready() def shutdown_preload(self) -> None: thread = self._preload_thread if thread is not None and thread.is_alive(): thread.join() self._preload_thread = None self._preload_result = None self._preload_error = None self._preload_chunk_id = None self._preload_done.clear() def get_preload_status(self) -> dict[str, Any]: return { "preload_chunk_id": -1 if self._preload_chunk_id is None else int(self._preload_chunk_id), "preload_ready": bool(self._preload_done.is_set()), "preload_running": bool(self._preload_thread is not None and self._preload_thread.is_alive()), "has_preload_result": bool(self._preload_result is not None), } def __len__(self) -> int: if self.current is None: return 0 return len(self.current["sample_times"]) def _cached_row(self, source: str, dat_idx: int) -> np.ndarray: if self.current is None: raise RuntimeError("RAM chunk is not loaded. Call load_chunk_sync first.") local_idx = self.current["local_maps"][source][int(dat_idx)] return self.current["arrays"][source][local_idx] def _load_input(self, name: str, sample_time: str) -> torch.Tensor: source = self.base._catalog_source_name(name) rows = [self._cached_row(source, self.base._idx_for(ts, source)) for ts in self.base._input_times(sample_time)] arr = np.stack(rows, axis=0) return torch.from_numpy(arr) def _load_label(self, label: str, sample_time: str) -> torch.Tensor: source = self.base._catalog_source_name(label) if label == "bt": rows = [self._cached_row(source, self.base._idx_for(ts, source)) for ts in self.base._bt_times(sample_time)] bt = np.stack(rows, axis=0).astype(np.float32, copy=False) if self.base.use_bt_mask: mask_source = self.base.bt_mask_source mask = self._cached_row(mask_source, self.base._idx_for(sample_time, mask_source)) bt = self.base.apply_bt_mask(bt, mask) return torch.from_numpy(bt.astype(np.float32, copy=False)) arr = self._cached_row(source, self.base._idx_for(sample_time, source)) return torch.from_numpy(arr) def __getitem__(self, idx: int) -> dict[str, Any]: if self.current is None: raise RuntimeError("RAM chunk is not loaded. Call load_chunk_sync first.") sample_time = self.current["sample_times"][int(idx)] return { "inputs": {name: self._load_input(name, sample_time) for name in self.base.required_inputs}, "labels": {name: self._load_label(name, sample_time) for name in self.base.required_labels}, "time": sample_time, } class FastUniqueRAMCachedDataset(Dataset): """ Full-split RAM cache for FastDataset without sliding-window duplication. It caches each required .dat row once, then reconstructs the original FastDataset sample dict on demand. """ def __init__( self, base: FastDataset, cache_dtype: str | np.dtype = "float16", read_block_rows: int = 64, verbose: bool = True, ): if not isinstance(base, FastDataset): raise TypeError("FastUniqueRAMCachedDataset only supports FastDataset") self.base = base self.cache_dtype = _numpy_dtype(str(cache_dtype)) self.read_block_rows = max(1, int(read_block_rows)) self.verbose = bool(verbose) refs: dict[str, set[int]] = {} for sample_time in self.base.samples: for source, indices in self._sample_refs(sample_time).items(): refs.setdefault(source, set()).update(indices) self.arrays: dict[str, np.ndarray] = {} self.local_maps: dict[str, dict[int, int]] = {} for source, idx_set in sorted(refs.items()): indices = np.array(sorted(idx_set), dtype=np.int64) self.arrays[source] = self._load_source_rows(source, indices) self.local_maps[source] = {int(idx): i for i, idx in enumerate(indices)} self.real_gb = sum(arr.nbytes for arr in self.arrays.values()) / (1024 ** 3) if self.verbose: detail = ", ".join(f"{src}={arr.shape}{arr.dtype}" for src, arr in self.arrays.items()) print( f"[FastUniqueRAMCachedDataset] samples={len(self.base)} " f"| {detail} | real={self.real_gb:.2f} GB" ) def _source_cache_dtype(self, source: str) -> np.dtype: store_dtype = self.base.stores[source].dtype if np.issubdtype(store_dtype, np.floating): return self.cache_dtype return store_dtype def _sample_refs(self, sample_time: str) -> dict[str, set[int]]: refs: dict[str, set[int]] = {} for name in self.base.required_inputs: source = self.base._catalog_source_name(name) refs.setdefault(source, set()) for ts in self.base._input_times(sample_time): refs[source].add(self.base._idx_for(ts, source)) for label in self.base.required_labels: source = self.base._catalog_source_name(label) refs.setdefault(source, set()) if label == "bt": if self.base.use_bt_mask: mask_source = self.base.bt_mask_source refs.setdefault(mask_source, set()).add(self.base._idx_for(sample_time, mask_source)) for ts in self.base._bt_times(sample_time): refs[source].add(self.base._idx_for(ts, source)) else: refs[source].add(self.base._idx_for(sample_time, source)) return refs def _load_source_rows(self, source: str, indices: np.ndarray) -> np.ndarray: store = self.base.stores[source] dtype = self._source_cache_dtype(source) out = np.empty((len(indices), *store.row_shape), dtype=dtype) for start in range(0, len(indices), self.read_block_rows): end = min(start + self.read_block_rows, len(indices)) out[start:end] = store._mm[indices[start:end]].astype(dtype, copy=False) return out def _cached_row(self, source: str, dat_idx: int) -> np.ndarray: local_idx = self.local_maps[source][int(dat_idx)] return self.arrays[source][local_idx] def _load_input(self, name: str, sample_time: str) -> torch.Tensor: source = self.base._catalog_source_name(name) rows = [self._cached_row(source, self.base._idx_for(ts, source)) for ts in self.base._input_times(sample_time)] return torch.from_numpy(np.stack(rows, axis=0)) def _load_label(self, label: str, sample_time: str) -> torch.Tensor: source = self.base._catalog_source_name(label) if label == "bt": rows = [self._cached_row(source, self.base._idx_for(ts, source)) for ts in self.base._bt_times(sample_time)] bt = np.stack(rows, axis=0).astype(np.float32, copy=False) if self.base.use_bt_mask: mask_source = self.base.bt_mask_source mask = self._cached_row(mask_source, self.base._idx_for(sample_time, mask_source)) bt = self.base.apply_bt_mask(bt, mask) return torch.from_numpy(bt.astype(np.float32, copy=False)) return torch.from_numpy(self._cached_row(source, self.base._idx_for(sample_time, source))) def __len__(self) -> int: return len(self.base.samples) def __getitem__(self, idx: int) -> dict[str, Any]: sample_time = self.base.samples[int(idx)] return { "inputs": {name: self._load_input(name, sample_time) for name in self.base.required_inputs}, "labels": {name: self._load_label(name, sample_time) for name in self.base.required_labels}, "time": sample_time, }