Download code/training/src/training_validation/ram_chunk.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 21 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/training_validation/ram_chunk.py
- Command line
-
hf download hf://lsh9034/ci-net/code/training/src/training_validation/ram_chunk.py
-
curl -L -o ram_chunk.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/training_validation/ram_chunk.py
21 kB
| 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") | |
| 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, | |
| } | |