poincare-hyper / src /env.py
DHDRL's picture
Rename env.py to src/env.py
84eb328 verified
Raw
History Blame Contribute Delete
6.54 kB
"""
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}