lsh9034's picture
Add files using upload-large-folder tool
76d61a0 verified
Raw History Blame Contribute Delete
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")
@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,
}