""" Multi-step self-supervised RL-style environment for Well-like data + Poincaré 8D. Note: SyntheticWellLike used to live in this file. Moved to synthetic_fields.py this session so that code needing only the plain data generator (no RL) doesn't pull in this module's gymnasium dependency. Re-imported below for backward compatibility with any code still doing `from .env import SyntheticWellLike`. """ from __future__ import annotations import numpy as np import torch from torch.utils.data import Dataset import gymnasium as gym from gymnasium import spaces from typing import Optional, Dict, Any from .normalization import FieldNormalizer from .synthetic_fields import SyntheticWellLike # noqa: F401 -- re-exported for compatibility class WellStreamDataset(Dataset): def __init__( self, dataset_name: str = "active_matter", split: str = "train", n_steps_input: int = 4, n_steps_output: int = 4, max_samples: int = 256, allow_synthetic_fallback: bool = False, ): from .provenance import DataLoadError self.max_samples = max_samples self.provenance = None try: from the_well.data import WellDataset print(f"[WellStream] Trying HF stream {dataset_name}/{split} ...") self.ds = WellDataset( well_base_path="hf://datasets/polymathic-ai/", well_dataset_name=dataset_name, well_split_name=split, n_steps_input=n_steps_input, n_steps_output=n_steps_output, ) self._len = min(len(self.ds), max_samples) self.provenance = "REAL_STREAMED" print(f"[WellStream] Real data OK – using {self._len} samples.") except Exception as e: if not allow_synthetic_fallback: raise DataLoadError( f"HF stream for '{dataset_name}/{split}' failed " f"({type(e).__name__}: {e}). Refusing to silently " f"substitute synthetic data. Pass " f"allow_synthetic_fallback=True if that is genuinely " f"what you want, or use get_synthetic_dataset() " f"explicitly.", outcome_code="STREAM_FAILED", ) from e print(f"[WellStream] WARNING: stream failed ({type(e).__name__}); " f"allow_synthetic_fallback=True — using SYNTHETIC data. " f"This dataset's provenance is 'SYNTHETIC', not real.") total_steps = n_steps_input + n_steps_output self.ds = SyntheticWellLike(n_samples=max_samples, n_steps=max(total_steps, 12)) self._len = len(self.ds) self.provenance = "SYNTHETIC" def __len__(self): return self._len def __getitem__(self, idx): return self.ds[idx] class MultiStepPoincareEnv(gym.Env): metadata = {"render_modes": []} def __init__( self, dataset: Dataset, normalizer: FieldNormalizer, encoder: torch.nn.Module, poincare_module, window: int = 4, horizon: int = 4, device: str = "cpu", ): super().__init__() self.dataset = dataset self.normalizer = normalizer self.encoder = encoder.to(device).eval() self.poincare = poincare_module self.window = window self.horizon = horizon self.device = device self.observation_space = spaces.Box(low=-np.inf, high=np.inf, shape=(8,), dtype=np.float32) self.action_space = spaces.Box(low=-1.5, high=1.5, shape=(8,), dtype=np.float32) self._traj = None self._t = 0 self._step_count = 0 self._current_latent = None def _load_traj(self, idx: int) -> torch.Tensor: item = self.dataset[idx % len(self.dataset)] if isinstance(item, dict): fields = None for k in ("fields", "input_fields", "x", "data"): if k in item and torch.is_tensor(item[k]): fields = item[k] break if fields is None: fields = next(v for v in item.values() if torch.is_tensor(v)) else: fields = item if fields.dim() == 5: fields = fields[0] fields = fields.float() frames = [self.normalizer.transform(fields[t]) for t in range(fields.shape[0])] return torch.stack(frames, dim=0).to(self.device) @torch.no_grad() def _encode(self, frames: torch.Tensor) -> torch.Tensor: x = frames[-1].unsqueeze(0) return self.encoder(x).squeeze(0) def reset(self, *, seed=None, options=None): super().reset(seed=seed) idx = np.random.randint(0, len(self.dataset)) self._traj = self._load_traj(idx) T = self._traj.shape[0] max_start = max(0, T - self.window - self.horizon) self._t = np.random.randint(0, max_start + 1) if max_start > 0 else 0 self._step_count = 0 window = self._traj[self._t : self._t + self.window] z = self._encode(window) self._current_latent = z return z.cpu().numpy().astype(np.float32), {"t": self._t, "traj_len": T} def step(self, action: np.ndarray): action_t = torch.as_tensor(action, device=self.device, dtype=torch.float32) pred_ball = self.poincare.expmap0(action_t.unsqueeze(0)).squeeze(0) next_t = self._t + self.window if next_t < self._traj.shape[0]: true_euc = self.encoder(self._traj[next_t].unsqueeze(0)).squeeze(0) true_ball = self.poincare.expmap0(true_euc.unsqueeze(0)).squeeze(0) else: true_ball = self.poincare.expmap0(self._current_latent.unsqueeze(0)).squeeze(0) dist = float(self.poincare.dist(pred_ball.unsqueeze(0), true_ball.unsqueeze(0)).item()) reward = -dist self._t += 1 self._step_count += 1 terminated = self._step_count >= self.horizon truncated = (self._t + self.window) >= self._traj.shape[0] if not (terminated or truncated): window = self._traj[self._t : self._t + self.window] z = self._encode(window) self._current_latent = z obs = z.cpu().numpy().astype(np.float32) else: obs = self._current_latent.cpu().numpy().astype(np.float32) return obs, reward, terminated, truncated, {"hyperbolic_dist": dist, "t": self._t}