{"role":"tool","content":[{"type":"tool-result","toolCallId":"call_PZrftVyq3ctEkxdO09Xxpi3F\nctc_05a1b235828e0c0e0169bba2b538d881969183fe6b3824bd46","toolName":"ApplyPatch","result":"The file /intern9/huhongkai/hs300_factor_lab/src/train_seq_gru_ddp_memmap.py has been updated.","experimental_content":[{"type":"text","text":"The file /intern9/huhongkai/hs300_factor_lab/src/train_seq_gru_ddp_memmap.py has been updated."}]}],"id":"call_PZrftVyq3ctEkxdO09Xxpi3F\nctc_05a1b235828e0c0e0169bba2b538d881969183fe6b3824bd46","providerOptions":{"cursor":{"highLevelToolCallResult":{"output":{"success":{"path":"/intern9/huhongkai/hs300_factor_lab/src/train_seq_gru_ddp_memmap.py","linesAdded":239,"linesRemoved":32,"diffString":" \n \n+@dataclass(frozen=True)\n+class MaterializedSeqBatch:\n+    xb: torch.Tensor\n+    y: torch.Tensor\n+    w: torch.Tensor\n+\n+\n class SeqMemmapBatchLoader:\n     def __init__(\n...\n \n \n+class MaterializeDevicePrefetchLoader:\n+    def __init__(\n+        self,\n+        loader,\n+        device: torch.device,\n+        seq_offsets: torch.Tensor,\n+        feat_mean: torch.Tensor,\n+        feat_std: torch.Tensor,\n+        prefetch_batches: int = 2,\n+        copy_non_blocking: bool = False,\n+    ):\n+        self.loader = loader\n+        self.device = device\n+        self.seq_offsets = seq_offsets\n+        self.feat_mean = feat_mean\n+        self.feat_std = feat_std\n+        self.prefetch_batches = max(1, int(prefetch_batches))\n+        self.copy_non_blocking = bool(copy_non_blocking)\n+\n+    def __len__(self) -> int:\n+        return len(self.loader)\n+\n+    def __getattr__(self, name: str):\n+        return getattr(self.loader, name)\n+\n+    def set_epoch(self, epoch: int) -> None:\n+        if hasattr(self.loader, \"set_epoch\"):\n+            self.loader.set_epoch(epoch)\n+\n+    def __iter__(self):\n+        if self.device.type != \"cuda\":\n+            for batch in self.loader:\n+                xb, yb, wb = materialize_batch(\n+                    batch,\n+                    seq_offsets=self.seq_offsets,\n+                    feat_mean=self.feat_mean,\n+                    feat_std=self.feat_std,\n+                    device=self.device,\n+                    copy_non_blocking=self.copy_non_blocking,\n+                )\n+                yield MaterializedSeqBatch(xb=xb, y=yb, w=wb)\n+            return\n+        result_q: queue.Queue = queue.Queue(maxsize=self.prefetch_batches)\n+        sentinel = object()\n+        device = self.device\n+        stop_event = threading.Event()\n+\n+        def put_result(item, event) -> bool:\n+            while not stop_event.is_set():\n+                try:\n+                    result_q.put((item, event), timeout=0.1)\n+                    return True\n+                except queue.Full:\n+                    continue\n+            return False\n+\n+        def worker():\n+            try:\n+                torch.cuda.set_device(device)\n+                stream = torch.cuda.Stream(device=device)\n+                for batch in self.loader:\n+                    if stop_event.is_set():\n+                        break\n+                    with torch.cuda.stream(stream):\n+                        xb, yb, wb = materialize_batch(\n+                            batch,\n+                            seq_offsets=self.seq_offsets,\n+                            feat_mean=self.feat_mean,\n+                            feat_std=self.feat_std,\n+                            device=device,\n+                            copy_non_blocking=self.copy_non_blocking,\n+                        )\n+                        event = torch.cuda.Event()\n+                        event.record(stream)\n+                    prefetched = MaterializedSeqBatch(xb=xb, y=yb, w=wb)\n+                    if not put_result(prefetched, event):\n+                        return\n+                put_result(sentinel, None)\n+            except Exception as exc:  # pragma: no cover - worker thread failure propagation\n+                put_result(exc, None)\n+\n+        thread = threading.Thread(target=worker, daemon=True)\n+        thread.start()\n+        current_stream = torch.cuda.current_stream(device)\n+        try:\n+            while True:\n+                item, event = result_q.get()\n+                if item is sentinel:\n+                    break\n+                if isinstance(item, Exception):\n+                    raise item\n+                current_stream.wait_event(event)\n+                item.xb.record_stream(current_stream)\n+                item.y.record_stream(current_stream)\n+                item.w.record_stream(current_stream)\n+                yield item\n+        finally:\n+            stop_event.set()\n+            thread.join()\n+\n+\n class ParallelCNN1D(nn.Module):\n     def __init__(self, in_channels: int, out_channels: int):\n...\n                 break\n             t1 = time.perf_counter()\n-            xb, yb, wb = materialize_batch(\n+            xb, yb, wb = resolve_model_batch(\n                 batch,\n                 seq_offsets=seq_offsets,\n...\n \n \n+def resolve_model_batch(\n+    batch,\n+    seq_offsets: torch.Tensor,\n+    feat_mean: torch.Tensor,\n+    feat_std: torch.Tensor,\n+    device: torch.device,\n+    copy_non_blocking: bool = False,\n+) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:\n+    if isinstance(batch, MaterializedSeqBatch):\n+        return batch.xb, batch.y, batch.w\n+    return materialize_batch(\n+        batch,\n+        seq_offsets=seq_offsets,\n+        feat_mean=feat_mean,\n+        feat_std=feat_std,\n+        device=device,\n+        copy_non_blocking=copy_non_blocking,\n+    )\n+\n+\n def sync_if_needed(device: torch.device) -> None:\n     if device.type == \"cuda\":\n...\n     train_iter = iter(loader)\n     warmup_collective(accelerator)\n+    if accelerator.device.type == \"cuda\":\n+        sync_if_needed(accelerator.device)\n+        torch.cuda.reset_peak_memory_stats(accelerator.device)\n     accum_steps = max(1, int(grad_accum_steps))\n     accum_micro = 0\n     optimizer.zero_grad(set_to_none=True)\n     for step_i in range(1, run_steps + 1):\n+        if accelerator.device.type == \"cuda\" and step_i == warmup + 1:\n+            sync_if_needed(accelerator.device)\n+            torch.cuda.reset_peak_memory_stats(accelerator.device)\n         t_step0 = time.perf_counter()\n         t0 = time.perf_counter()\n...\n         should_step = accum_micro >= accum_steps or last_batch\n         group_total = min(accum_steps, accum_micro + max(0, run_steps - step_i))\n-        xb, yb, wb = materialize_batch(\n+        xb, yb, wb = resolve_model_batch(\n             batch,\n             seq_offsets=seq_offsets,\n...\n         if accelerator.is_main_process and (step_i == run_steps or step_i % max(1, run_steps // 4) == 0):\n             print(f\"[profile-train] step {step_i}/{run_steps} loss={loss_raw.detach().float().item():.6f}\", flush=True)\n+    peak_allocated_bytes = 0.0\n+    peak_reserved_bytes = 0.0\n+    if accelerator.device.type == \"cuda\":\n+        sync_if_needed(accelerator.device)\n+        peak_allocated_bytes = float(torch.cuda.max_memory_allocated(accelerator.device))\n+        peak_reserved_bytes = float(torch.cuda.max_memory_reserved(accelerator.device))\n     global_batch = int(loader.batch_size) * int(accelerator.num_processes)\n     summary = summarize_timing_across_ranks(\n...\n     summary[\"warmup_steps\"] = warmup\n     summary[\"profile_steps\"] = active\n+    if accelerator.device.type == \"cuda\":\n+        local_mem = torch.tensor(\n+            [peak_allocated_bytes, peak_reserved_bytes],\n+            dtype=torch.float64,\n+            device=accelerator.device,\n+        )[None, :]\n+        gathered_mem = accelerator.gather(local_mem)\n+        if accelerator.is_main_process:\n+            gathered_mem_np = gathered_mem.detach().cpu().numpy()\n+            peak_allocated_max = 0.0\n+            peak_reserved_max = 0.0\n+            for rank_i, row in enumerate(gathered_mem_np):\n+                allocated_bytes = float(row[0])\n+                reserved_bytes = float(row[1])\n+                peak_allocated_max = max(peak_allocated_max, allocated_bytes)\n+                peak_reserved_max = max(peak_reserved_max, reserved_bytes)\n+                if rank_i < len(summary.get(\"per_rank\", [])):\n+                    summary[\"per_rank\"][rank_i][\"peak_allocated_bytes\"] = allocated_bytes\n+                    summary[\"per_rank\"][rank_i][\"peak_reserved_bytes\"] = reserved_bytes\n+                    summary[\"per_rank\"][rank_i][\"peak_allocated_gib\"] = allocated_bytes / float(1024**3)\n+                    summary[\"per_rank\"][rank_i][\"peak_reserved_gib\"] = reserved_bytes / float(1024**3)\n+            summary[\"peak_allocated_bytes_max\"] = peak_allocated_max\n+            summary[\"peak_reserved_bytes_max\"] = peak_reserved_max\n+            summary[\"peak_allocated_gib_max\"] = peak_allocated_max / float(1024**3)\n+            summary[\"peak_reserved_gib_max\"] = peak_reserved_max / float(1024**3)\n     return summary\n \n...\n         else None\n     )\n-    if device_transfer_mode == \"thread_prefetch\":\n-        train_loader = DeviceTransferPrefetchLoader(\n-            train_loader,\n-            device=accelerator.device,\n-            prefetch_batches=device_transfer_prefetch_batches,\n-            copy_non_blocking=use_pinned_transfer,\n-        )\n-        if train_eval_loader is not None:\n-            train_eval_loader = DeviceTransferPrefetchLoader(\n-                train_eval_loader,\n-                device=accelerator.device,\n-                prefetch_batches=device_transfer_prefetch_batches,\n-                copy_non_blocking=use_pinned_transfer,\n-            )\n-        if valid_loader is not None:\n-            valid_loader = DeviceTransferPrefetchLoader(\n-                valid_loader,\n-                device=accelerator.device,\n-                prefetch_batches=device_transfer_prefetch_batches,\n-                copy_non_blocking=use_pinned_transfer,\n-            )\n-        if test_loader is not None:\n-            test_loader = DeviceTransferPrefetchLoader(\n-                test_loader,\n-                device=accelerator.device,\n-                prefetch_batches=device_transfer_prefetch_batches,\n-                copy_non_blocking=use_pinned_transfer,\n-            )\n     phase_timings[\"test_loader_init_sec\"] = float(time.perf_counter() - loader_init_t0)\n     stats_stride = max(20, args.sample_stride)\n...\n     seq_offsets = torch.arange(-(args.seq_len - 1), 1, dtype=torch.long, device=accelerator.device)\n     phase_timings[\"device_tensor_init_sec\"] = float(time.perf_counter() - tensor_init_t0)\n+    device_prefetch_t0 = time.perf_counter()\n+    if device_transfer_mode == \"thread_prefetch\":\n+        train_loader = DeviceTransferPrefetchLoader(\n+            train_loader,\n+            device=accelerator.device,\n+            prefetch_batches=device_transfer_prefetch_batches,\n+            copy_non_blocking=use_pinned_transfer,\n+        )\n+        if train_eval_loader is not None:\n+            train_eval_loader = DeviceTransferPrefetchLoader(\n+                train_eval_loader,\n+                device=accelerator.device,\n+                prefetch_batches=device_transfer_prefetch_batches,\n+                copy_non_blocking=use_pinned_transfer,\n+            )\n+        if valid_loader is not None:\n+            valid_loader = DeviceTransferPrefetchLoader(\n+                valid_loader,\n+                device=accelerator.device,\n+                prefetch_batches=device_transfer_prefetch_batches,\n+                copy_non_blocking=use_pinned_transfer,\n+            )\n+        if test_loader is not None:\n+            test_loader = DeviceTransferPrefetchLoader(\n+                test_loader,\n+                device=accelerator.device,\n+                prefetch_batches=device_transfer_prefetch_batches,\n+                copy_non_blocking=use_pinned_transfer,\n+            )\n+    elif device_transfer_mode == \"materialize_prefetch\":\n+        train_loader = MaterializeDevicePrefetchLoader(\n+            train_loader,\n+            device=accelerator.device,\n+            seq_offsets=seq_offsets,\n+            feat_mean=feat_mean,\n+            feat_std=feat_std,\n+            prefetch_batches=device_transfer_prefetch_batches,\n+            copy_non_blocking=use_pinned_transfer,\n+        )\n+        if train_eval_loader is not None:\n+            train_eval_loader = MaterializeDevicePrefetchLoader(\n+                train_eval_loader,\n+                device=accelerator.device,\n+                seq_offsets=seq_offsets,\n+                feat_mean=feat_mean,\n+                feat_std=feat_std,\n+                prefetch_batches=device_transfer_prefetch_batches,\n+                copy_non_blocking=use_pinned_transfer,\n+            )\n+        if valid_loader is not None:\n+            valid_loader = MaterializeDevicePrefetchLoader(\n+                valid_loader,\n+                device=accelerator.device,\n+                seq_offsets=seq_offsets,\n+                feat_mean=feat_mean,\n+                feat_std=feat_std,\n+                prefetch_batches=device_transfer_prefetch_batches,\n+                copy_non_blocking=use_pinned_transfer,\n+            )\n+        if test_loader is not None:\n+            test_loader = MaterializeDevicePrefetchLoader(\n+                test_loader,\n+                device=accelerator.device,\n+                seq_offsets=seq_offsets,\n+                feat_mean=feat_mean,\n+                feat_std=feat_std,\n+                prefetch_batches=device_transfer_prefetch_batches,\n+                copy_non_blocking=use_pinned_transfer,\n+            )\n+    phase_timings[\"device_prefetch_wrap_sec\"] = float(time.perf_counter() - device_prefetch_t0)\n \n     output_init_t0 = time.perf_counter()\n...\n             group_total = min(grad_accum_steps, accum_micro + max(0, total_train_batches - batch_i))\n             measure_first_step = bool(profile_phases and first_train_step_local is None and epoch == 1)\n-            xb, yb, wb = materialize_batch(\n+            xb, yb, wb = resolve_model_batch(\n                 batch,\n                 seq_offsets=seq_offsets,\n...\n         type=str,\n         default=\"thread_prefetch\",\n-        choices=[\"direct\", \"thread_prefetch\"],\n+        choices=[\"direct\", \"thread_prefetch\", \"materialize_prefetch\"],\n     )\n     p.add_argument(\"--device-transfer-prefetch-batches\", type=int, default=2)","beforeFullFileContent":"import argparse\nimport csv\nimport hashlib\nimport json\nimport math\nimport os\nimport queue\nimport random\nimport threading\nimport time\nfrom concurrent.futures import Future, ThreadPoolExecutor\nfrom contextlib import nullcontext\nfrom dataclasses import dataclass\nfrom datetime import timedelta\nfrom pathlib import Path\nfrom typing import Dict, List, Tuple\n\nimport numpy as np\nimport torch\nimport torch.distributed as dist\nimport torch.nn as nn\nfrom torch.nn.parallel import DistributedDataParallel as DDP\nfrom accelerate import Accelerator\n\nfrom common import ensure_dir, load_config, load_split_view_meta, resolve_named_day_dirs, resolve_row_root\n\n\ndef set_global_seed(seed: int) -> None:\n    seed = int(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n\n\ndef build_epoch_lr_schedule(train_cfg: dict, epochs: int) -> List[float]:\n    base_lr = float(train_cfg[\"learning_rate\"])\n    raw_values = train_cfg.get(\"lr_epoch_values\")\n    if raw_values is None:\n        return [base_lr for _ in range(int(epochs))]\n    values = [float(x) for x in raw_values]\n    if len(values) != int(epochs):\n        raise ValueError(\n            f\"train.lr_epoch_values length ({len(values)}) must match epochs ({int(epochs)}).\"\n        )\n    return values\n\n\ndef set_optimizer_lr(optimizer, lr: float) -> None:\n    lr = float(lr)\n    for group in optimizer.param_groups:\n        group[\"lr\"] = lr\n\n\nclass AccelerateRuntime:\n    def __init__(self, use_amp: bool):\n        self.accelerator = Accelerator(mixed_precision=\"fp16\" if bool(use_amp) else \"no\")\n        self.device = self.accelerator.device\n        self.process_index = int(self.accelerator.process_index)\n        self.num_processes = int(self.accelerator.num_processes)\n        self.is_main_process = bool(self.accelerator.is_main_process)\n\n    def autocast(self):\n        return self.accelerator.autocast()\n\n    def backward(self, loss: torch.Tensor) -> None:\n        self.accelerator.backward(loss)\n\n    def step_optimizer(self, optimizer) -> None:\n        optimizer.step()\n\n    def prepare(self, model, optimizer):\n        return self.accelerator.prepare(model, optimizer)\n\n    def unwrap_model(self, model):\n        return self.accelerator.unwrap_model(model)\n\n    def reduce(self, tensor: torch.Tensor, reduction: str = \"sum\") -> torch.Tensor:\n        return self.accelerator.reduce(tensor, reduction=reduction)\n\n    def gather(self, tensor: torch.Tensor) -> torch.Tensor:\n        return self.accelerator.gather(tensor)\n\n    def wait_for_everyone(self) -> None:\n        self.accelerator.wait_for_everyone()\n\n    def close(self) -> None:\n        return None\n\n\nclass NativeDDPRuntime:\n    def __init__(self, use_amp: bool, enable_static_graph: bool = True):\n        if torch.cuda.is_available():\n            local_rank = int(os.environ.get(\"LOCAL_RANK\", 0))\n            torch.cuda.set_device(local_rank)\n            self.device = torch.device(\"cuda\", local_rank)\n        else:\n            self.device = torch.device(\"cpu\")\n        requested_world = int(os.environ.get(\"WORLD_SIZE\", \"1\"))\n        self._distributed = requested_world > 1\n        if self._distributed and not dist.is_initialized():\n            backend = \"nccl\" if self.device.type == \"cuda\" else \"gloo\"\n            dist.init_process_group(backend=backend, timeout=timedelta(seconds=7200))\n        if dist.is_initialized():\n            self.process_index = int(dist.get_rank())\n            self.num_processes = int(dist.get_world_size())\n        else:\n            self.process_index = 0\n            self.num_processes = 1\n        self.is_main_process = self.process_index == 0\n        self.use_amp = bool(use_amp) and self.device.type == \"cuda\"\n        self.enable_static_graph = bool(enable_static_graph)\n        self.scaler = torch.cuda.amp.GradScaler(enabled=self.use_amp)\n\n    def autocast(self):\n        if self.device.type == \"cuda\":\n            return torch.autocast(device_type=\"cuda\", dtype=torch.float16, enabled=self.use_amp)\n        return nullcontext()\n\n    def backward(self, loss: torch.Tensor) -> None:\n        if self.use_amp:\n            self.scaler.scale(loss).backward()\n        else:\n            loss.backward()\n\n    def step_optimizer(self, optimizer) -> None:\n        if self.use_amp:\n            self.scaler.step(optimizer)\n            self.scaler.update()\n        else:\n            optimizer.step()\n\n    def prepare(self, model, optimizer):\n        model = model.to(self.device)\n        if self.num_processes > 1:\n            ddp_kwargs = {\n                \"broadcast_buffers\": False,\n                \"gradient_as_bucket_view\": True,\n            }\n            if self.device.type == \"cuda\":\n                ddp_kwargs[\"device_ids\"] = [self.device.index]\n                ddp_kwargs[\"output_device\"] = self.device.index\n            if self.enable_static_graph:\n                try:\n                    model = DDP(model, static_graph=True, **ddp_kwargs)\n                except TypeError:\n                    model = DDP(model, **ddp_kwargs)\n            else:\n                model = DDP(model, **ddp_kwargs)\n        return model, optimizer\n\n    def unwrap_model(self, model):\n        return model.module if isinstance(model, DDP) else model\n\n    def reduce(self, tensor: torch.Tensor, reduction: str = \"sum\") -> torch.Tensor:\n        if self.num_processes <= 1:\n            return tensor\n        out = tensor.clone()\n        reduction_name = str(reduction).lower()\n        if reduction_name == \"sum\":\n            op = dist.ReduceOp.SUM\n            dist.all_reduce(out, op=op)\n        elif reduction_name == \"mean\":\n            dist.all_reduce(out, op=dist.ReduceOp.SUM)\n            out = out / float(self.num_processes)\n        elif reduction_name == \"max\":\n            dist.all_reduce(out, op=dist.ReduceOp.MAX)\n        else:\n            raise ValueError(f\"Unsupported reduction: {reduction}\")\n        return out\n\n    def gather(self, tensor: torch.Tensor) -> torch.Tensor:\n        if self.num_processes <= 1:\n            return tensor\n        gather_list = [torch.empty_like(tensor) for _ in range(self.num_processes)]\n        dist.all_gather(gather_list, tensor)\n        return torch.cat(gather_list, dim=0)\n\n    def wait_for_everyone(self) -> None:\n        if self.num_processes > 1:\n            if self.device.type == \"cuda\":\n                dist.barrier(device_ids=[self.device.index])\n            else:\n                dist.barrier()\n\n    def close(self) -> None:\n        if dist.is_initialized():\n            dist.destroy_process_group()\n\n\ndef build_runtime(runtime_backend: str, use_amp: bool, enable_static_graph: bool = True):\n    backend = str(runtime_backend).strip().lower()\n    if backend == \"native\":\n        return NativeDDPRuntime(use_amp=use_amp, enable_static_graph=enable_static_graph)\n    if backend == \"accelerate\":\n        return AccelerateRuntime(use_amp=use_amp)\n    raise ValueError(f\"Unsupported runtime_backend: {runtime_backend}\")\n\n\n@dataclass\nclass DayStore:\n    name: str\n    n_rows: int\n    n_factors: int\n    x: np.memmap\n    y: np.memmap\n    w: np.memmap\n    sym_start: np.memmap\n    sym_end: np.memmap\n    need_clean: bool\n\n\n_SPLIT_X_ARRAY_CACHE: Dict[str, np.ndarray] = {}\n\n\ndef _resolve_source_x_path(day_dir: Path, meta: dict) -> Path | None:\n    explicit = str(meta.get(\"source_x_path\") or \"\").strip()\n    if explicit:\n        return Path(explicit)\n    source_root = str(meta.get(\"source_root\") or \"\").strip()\n    source_split = str(meta.get(\"source_split\") or \"\").strip()\n    if source_root and source_split:\n        return Path(source_root) / source_split / \"x.npy\"\n    return None\n\n\ndef _load_cached_x_npy(path: Path) -> np.ndarray:\n    key = str(path.resolve())\n    arr = _SPLIT_X_ARRAY_CACHE.get(key)\n    if arr is None:\n        arr = np.load(path, mmap_mode=\"r\", allow_pickle=False)\n        _SPLIT_X_ARRAY_CACHE[key] = arr\n    return arr\n\n\ndef list_ready_days(row_root: Path) -> List[Path]:\n    days = []\n    for p in sorted(row_root.iterdir()):\n        if p.is_dir() and (p / \"_SUCCESS\").exists():\n            days.append(p)\n    return days\n\n\ndef split_days(days: List[Path], ratio: float) -> Tuple[List[Path], List[Path]]:\n    n = len(days)\n    if n < 2:\n        return days, days\n    cut = max(1, min(n - 1, int(n * ratio)))\n    return days[:cut], days[cut:]\n\n\ndef _day_str_from_dir(day_dir: Path) -> str:\n    name = day_dir.name\n    if len(name) >= 8:\n        day = name[:8]\n        if day.isdigit():\n            return day\n    return \"\"\n\n\ndef filter_days_by_date(days: List[Path], start_date: str, end_date: str) -> List[Path]:\n    s = (start_date or \"\").strip()\n    e = (end_date or \"\").strip()\n    if not s and not e:\n        return list(days)\n    if s and (len(s) != 8 or not s.isdigit()):\n        raise ValueError(f\"Invalid start_date: {s}\")\n    if e and (len(e) != 8 or not e.isdigit()):\n        raise ValueError(f\"Invalid end_date: {e}\")\n    lo = s if s else \"00000000\"\n    hi = e if e else \"99999999\"\n    if lo > hi:\n        raise ValueError(f\"Invalid date window: {lo} > {hi}\")\n    out: List[Path] = []\n    for d in days:\n        day = _day_str_from_dir(d)\n        if day and lo <= day <= hi:\n            out.append(d)\n    return out\n\n\ndef limit_days(days: List[Path], day_limit: int) -> List[Path]:\n    if day_limit <= 0 or day_limit >= len(days):\n        return days\n    return days[:day_limit]\n\n\ndef load_day_meta(day_dir: Path) -> dict:\n    return json.loads((day_dir / \"meta.json\").read_text(encoding=\"utf-8\"))\n\n\ndef load_day_store(day_dir: Path, prefer_fp16: bool) -> DayStore:\n    meta = load_day_meta(day_dir)\n    n_rows = int(meta[\"n_rows\"])\n    n_factors = int(meta[\"n_factors\"])\n    source_x_path = _resolve_source_x_path(day_dir, meta)\n    if source_x_path is not None:\n        row_start = int(meta.get(\"source_row_start\", 0))\n        row_stop = int(meta.get(\"source_row_stop\", row_start + n_rows))\n        x_all = _load_cached_x_npy(source_x_path)\n        x = x_all[row_start:row_stop]\n        need_clean = True\n    else:\n        fp16_path = day_dir / \"x_top200_f16_filled.memmap\"\n        if prefer_fp16 and fp16_path.exists():\n            x = np.memmap(fp16_path, mode=\"r\", dtype=np.float16, shape=(n_rows, n_factors))\n            need_clean = False\n        else:\n            x = np.memmap(day_dir / \"x_top200_float32.memmap\", mode=\"r\", dtype=np.float32, shape=(n_rows, n_factors))\n            need_clean = True\n    y = np.memmap(day_dir / \"y_sum_float32.memmap\", mode=\"r\", dtype=np.float32, shape=(n_rows,))\n    w = np.memmap(day_dir / \"w_float32.memmap\", mode=\"r\", dtype=np.float32, shape=(n_rows,))\n    sym_n = int(meta[\"symbol_count_meta\"])\n    sym_start = np.memmap(day_dir / \"symbol_start_idx_int64.memmap\", mode=\"r\", dtype=np.int64, shape=(sym_n,))\n    sym_end = np.memmap(day_dir / \"symbol_end_idx_int64.memmap\", mode=\"r\", dtype=np.int64, shape=(sym_n,))\n    return DayStore(day_dir.name, n_rows, n_factors, x, y, w, sym_start, sym_end, need_clean)\n\n\ndef load_day_symbol_bounds(day_dir: Path) -> Tuple[np.memmap, np.memmap]:\n    meta = load_day_meta(day_dir)\n    sym_n = int(meta[\"symbol_count_meta\"])\n    sym_start = np.memmap(day_dir / \"symbol_start_idx_int64.memmap\", mode=\"r\", dtype=np.int64, shape=(sym_n,))\n    sym_end = np.memmap(day_dir / \"symbol_end_idx_int64.memmap\", mode=\"r\", dtype=np.int64, shape=(sym_n,))\n    return sym_start, sym_end\n\n\ndef build_valid_end_indices_from_bounds(\n    sym_start: np.ndarray,\n    sym_end: np.ndarray,\n    seq_len: int,\n    sample_stride: int,\n) -> np.ndarray:\n    stride = int(sample_stride)\n    starts = np.asarray(sym_start, dtype=np.int64) + int(seq_len) - 1\n    # symbol_end_idx is exclusive in this memmap layout.\n    ends = np.asarray(sym_end, dtype=np.int64) - 1\n    valid_mask = starts <= ends\n    if not np.any(valid_mask):\n        return np.empty((0,), dtype=np.int64)\n    starts = starts[valid_mask]\n    ends = ends[valid_mask]\n    lengths = ((ends - starts) // stride + 1).astype(np.int64, copy=False)\n    total = int(lengths.sum())\n    if total <= 0:\n        return np.empty((0,), dtype=np.int64)\n    repeated_starts = np.repeat(starts, lengths)\n    group_offsets = np.repeat(np.cumsum(lengths, dtype=np.int64) - lengths, lengths)\n    intra_offsets = np.arange(total, dtype=np.int64) - group_offsets\n    return repeated_starts + intra_offsets * stride\n\n\ndef build_valid_end_indices(store: DayStore, seq_len: int, sample_stride: int) -> np.ndarray:\n    return build_valid_end_indices_from_bounds(\n        store.sym_start,\n        store.sym_end,\n        seq_len=seq_len,\n        sample_stride=sample_stride,\n    )\n\n\ndef build_end_index_cache_path(cache_dir: Path, day_dir: Path, seq_len: int, sample_stride: int) -> Path:\n    key_src = \"\\n\".join(\n        [\n            str(day_dir.parent),\n            str(day_dir.name),\n            str(int(seq_len)),\n            str(int(sample_stride)),\n        ]\n    )\n    key = hashlib.sha1(key_src.encode(\"utf-8\")).hexdigest()[:16]\n    return cache_dir / f\"{day_dir.name}_seq{int(seq_len)}_stride{int(sample_stride)}_{key}.npy\"\n\n\ndef load_end_index_cache(path: Path) -> np.ndarray:\n    arr = np.load(path, allow_pickle=False)\n    return np.asarray(arr, dtype=np.int64)\n\n\ndef wait_for_end_index_cache(path: Path, timeout_sec: float = 600.0) -> np.ndarray:\n    deadline = time.time() + float(timeout_sec)\n    last_err = None\n    while time.time() < deadline:\n        if path.exists() and path.stat().st_size > 0:\n            try:\n                return load_end_index_cache(path)\n            except Exception as exc:  # pragma: no cover - transient partial-write case\n                last_err = exc\n        time.sleep(0.2)\n    if last_err is not None:\n        raise TimeoutError(f\"Timed out waiting for end-index cache {path}: {last_err}\") from last_err\n    raise TimeoutError(f\"Timed out waiting for end-index cache {path}\")\n\n\ndef load_or_build_valid_end_indices(\n    day_dir: Path,\n    seq_len: int,\n    sample_stride: int,\n    cache_dir: Path | None,\n    cache_writer: bool,\n) -> np.ndarray:\n    if cache_dir is None:\n        sym_start, sym_end = load_day_symbol_bounds(day_dir)\n        return build_valid_end_indices_from_bounds(sym_start, sym_end, seq_len=seq_len, sample_stride=sample_stride)\n    cache_path = build_end_index_cache_path(cache_dir, day_dir, seq_len=seq_len, sample_stride=sample_stride)\n    if cache_path.exists() and cache_path.stat().st_size > 0:\n        return load_end_index_cache(cache_path)\n    if not cache_writer:\n        return wait_for_end_index_cache(cache_path)\n    sym_start, sym_end = load_day_symbol_bounds(day_dir)\n    ends = build_valid_end_indices_from_bounds(sym_start, sym_end, seq_len=seq_len, sample_stride=sample_stride)\n    tmp_path = cache_path.with_suffix(f\".tmp.{int(time.time() * 1000)}.{os.getpid()}.npy\")\n    np.save(tmp_path, ends)\n    tmp_path.replace(cache_path)\n    return ends\n\n\ndef truncate_end_indices(end_indices: List[np.ndarray], max_samples: int) -> List[np.ndarray]:\n    if max_samples <= 0:\n        return end_indices\n    remaining = int(max_samples)\n    trimmed: List[np.ndarray] = []\n    for ends in end_indices:\n        if remaining <= 0:\n            trimmed.append(np.empty((0,), dtype=np.int64))\n            continue\n        take = min(int(ends.shape[0]), remaining)\n        trimmed.append(ends[:take])\n        remaining -= take\n    return trimmed\n\n\ndef split_contiguous_even(total: int, rank: int, world_size: int) -> Tuple[int, int]:\n    if total <= 0:\n        return 0, 0\n    start = (total * rank) // max(1, world_size)\n    end = (total * (rank + 1)) // max(1, world_size)\n    return int(start), int(end)\n\n\n@dataclass(frozen=True)\nclass BatchSlice:\n    day_i: int\n    start: int\n    stop: int\n\n\n@dataclass(frozen=True)\nclass BatchPlanItem:\n    parts: Tuple[BatchSlice, ...]\n    pad_size: int = 0\n\n\n@dataclass(frozen=True)\nclass PackedSeqBatchPart:\n    span_x: torch.Tensor\n    local_end: torch.Tensor\n    y: torch.Tensor\n    w: torch.Tensor\n\n\n@dataclass(frozen=True)\nclass PackedSeqBatch:\n    parts: Tuple[PackedSeqBatchPart, ...]\n\n\nclass SeqMemmapBatchLoader:\n    def __init__(\n        self,\n        day_dirs: List[Path],\n        seq_len: int,\n        sample_stride: int,\n        batch_size: int,\n        shuffle: bool,\n        max_samples: int = -1,\n        prefer_fp16: bool = True,\n        rank: int = 0,\n        world_size: int = 1,\n        seed: int = 20260312,\n        pin_memory: bool = True,\n        loader_threads: int = 1,\n        prefetch_batches: int = 1,\n        pad_last_batch: bool = False,\n        batch_overlap: int = 0,\n        index_cache_dir: Path | None = None,\n    ):\n        self.seq_len = seq_len\n        self.batch_size = int(batch_size)\n        self.batch_overlap = max(0, int(batch_overlap))\n        if self.batch_overlap >= self.batch_size:\n            raise ValueError(\n                f\"batch_overlap must be smaller than batch_size, got overlap={self.batch_overlap} batch_size={self.batch_size}\"\n            )\n        self.batch_stride = self.batch_size - self.batch_overlap\n        self.shuffle = bool(shuffle)\n        self.rank = int(rank)\n        self.world_size = int(world_size)\n        self.seed = int(seed)\n        self.day_dirs = list(day_dirs)\n        self.prefer_fp16 = bool(prefer_fp16)\n        self.pin_memory = bool(pin_memory) and torch.cuda.is_available()\n        self.loader_threads = max(1, int(loader_threads))\n        self.prefetch_batches = max(1, int(prefetch_batches))\n        self.pad_last_batch = bool(pad_last_batch)\n        self.epoch = 0\n        self.index_cache_dir = index_cache_dir\n        self.stores: List[DayStore | None] = [None for _ in self.day_dirs]\n        end_indices = [\n            load_or_build_valid_end_indices(\n                d,\n                seq_len=seq_len,\n                sample_stride=sample_stride,\n                cache_dir=self.index_cache_dir,\n                cache_writer=(self.rank == 0),\n            )\n            for d in self.day_dirs\n        ]\n        self.end_indices: List[np.ndarray] = truncate_end_indices(end_indices, max_samples=max_samples)\n        self.total_samples = int(sum(int(ends.shape[0]) for ends in self.end_indices))\n        self.seq_offsets = np.arange(-(self.seq_len - 1), 1, dtype=np.int64)\n        self.rank_sample_counts = self._compute_rank_sample_counts()\n        self.rank_total_samples = [int(sum(day_counts)) for day_counts in self.rank_sample_counts]\n        self.rank_unique_batches = [\n            self._compute_batch_count(total) for total in self.rank_total_samples\n        ]\n        self.total_batches = max(self.rank_unique_batches, default=0)\n\n    def __len__(self) -> int:\n        return self.total_batches\n\n    def set_epoch(self, epoch: int) -> None:\n        self.epoch = int(epoch)\n\n    def _get_store(self, day_i: int) -> DayStore:\n        store = self.stores[day_i]\n        if store is None:\n            store = load_day_store(self.day_dirs[day_i], prefer_fp16=self.prefer_fp16)\n            self.stores[day_i] = store\n        return store\n\n    def _compute_rank_sample_counts(self) -> List[List[int]]:\n        counts: List[List[int]] = [[] for _ in range(self.world_size)]\n        for ends in self.end_indices:\n            n = int(ends.shape[0])\n            for rank in range(self.world_size):\n                start, stop = split_contiguous_even(n, rank, self.world_size)\n                counts[rank].append(max(0, stop - start))\n        return counts\n\n    def _compute_batch_count(self, total: int) -> int:\n        total = int(total)\n        if total <= 0:\n            return 0\n        if total <= self.batch_size:\n            return 1\n        return 1 + int(math.ceil((total - self.batch_size) / self.batch_stride))\n\n    def _tail_parts(self, parts: List[BatchSlice], keep: int) -> List[BatchSlice]:\n        keep = max(0, int(keep))\n        if keep <= 0:\n            return []\n        out: List[BatchSlice] = []\n        remaining = keep\n        for part in reversed(parts):\n            part_len = int(part.stop - part.start)\n            if part_len <= 0:\n                continue\n            take = min(remaining, part_len)\n            out.append(BatchSlice(day_i=part.day_i, start=part.stop - take, stop=part.stop))\n            remaining -= take\n            if remaining == 0:\n                break\n        if remaining != 0:\n            raise RuntimeError(f\"Failed to preserve batch overlap keep={keep}, remaining={remaining}\")\n        out.reverse()\n        return out\n\n    def _build_rank_plan(self) -> List[BatchPlanItem]:\n        # Rotate the contiguous shard every epoch so each rank sees different regions over time.\n        shard_rank = (self.rank + self.epoch) % max(1, self.world_size) if self.shuffle else self.rank\n        items: List[BatchPlanItem] = []\n        current_parts: List[BatchSlice] = []\n        filled = 0\n        for day_i, ends in enumerate(self.end_indices):\n            total = int(ends.shape[0])\n            start, stop = split_contiguous_even(total, shard_rank, self.world_size)\n            if stop <= start:\n                continue\n            cursor = int(start)\n            stop = int(stop)\n            while cursor < stop:\n                need = self.batch_size - filled\n                take = min(need, stop - cursor)\n                current_parts.append(BatchSlice(day_i=day_i, start=cursor, stop=cursor + take))\n                cursor += take\n                filled += take\n                if filled == self.batch_size:\n                    emitted_parts = tuple(current_parts)\n                    items.append(BatchPlanItem(parts=emitted_parts, pad_size=0))\n                    current_parts = self._tail_parts(list(emitted_parts), self.batch_overlap)\n                    filled = self.batch_overlap\n        if current_parts:\n            pad_size = self.batch_size - filled if self.pad_last_batch else 0\n            items.append(BatchPlanItem(parts=tuple(current_parts), pad_size=pad_size))\n        if not items:\n            return []\n        if self.shuffle and len(items) > 1:\n            rng = np.random.default_rng(self.seed + self.epoch)\n            if self.batch_overlap > 0:\n                shift = int(rng.integers(len(items)))\n                if shift > 0:\n                    items = items[shift:] + items[:shift]\n            else:\n                order = rng.permutation(len(items))\n                items = [items[int(i)] for i in order.tolist()]\n        if len(items) < self.total_batches:\n            base = list(items)\n            pad_idx = 0\n            while len(items) < self.total_batches:\n                items.append(base[pad_idx % len(base)])\n                pad_idx += 1\n        return items\n\n    def _load_batch_part(self, part: BatchSlice) -> PackedSeqBatchPart:\n        day_i, start, stop = part.day_i, part.start, part.stop\n        store = self._get_store(day_i)\n        batch_end = self.end_indices[day_i][start:stop]\n        if batch_end.size == 0:\n            raise RuntimeError(f\"Empty batch slice for day_i={day_i}, start={start}, stop={stop}\")\n        span_start = int(batch_end[0]) - self.seq_len + 1\n        span_end = int(batch_end[-1])\n        span_x_dtype = np.float32 if store.need_clean else store.x.dtype\n        span_x = np.array(store.x[span_start : span_end + 1], dtype=span_x_dtype, copy=True)\n        if store.need_clean:\n            np.nan_to_num(span_x, copy=False, nan=0.0, posinf=0.0, neginf=0.0)\n        local_end = np.ascontiguousarray(batch_end - span_start, dtype=np.int64)\n        y = np.array(store.y[batch_end], dtype=np.float32, copy=True)\n        w = np.array(store.w[batch_end], dtype=np.float32, copy=True)\n        span_xb = torch.from_numpy(np.ascontiguousarray(span_x))\n        local_endb = torch.from_numpy(local_end)\n        yb = torch.from_numpy(np.ascontiguousarray(y))\n        wb = torch.from_numpy(np.ascontiguousarray(w))\n        if self.pin_memory:\n            span_xb = span_xb.pin_memory()\n            local_endb = local_endb.pin_memory()\n            yb = yb.pin_memory()\n            wb = wb.pin_memory()\n        return PackedSeqBatchPart(span_x=span_xb, local_end=local_endb, y=yb, w=wb)\n\n    def _pad_batch_part(self, part: PackedSeqBatchPart, pad_size: int) -> PackedSeqBatchPart:\n        if pad_size <= 0:\n            return part\n        local_end = torch.cat([part.local_end, part.local_end[-1:].repeat(int(pad_size))], dim=0)\n        y = torch.cat([part.y, part.y[-1:].repeat(int(pad_size))], dim=0)\n        w = torch.cat([part.w, part.w[-1:].repeat(int(pad_size))], dim=0)\n        if self.pin_memory:\n            local_end = local_end.pin_memory()\n            y = y.pin_memory()\n            w = w.pin_memory()\n        return PackedSeqBatchPart(span_x=part.span_x, local_end=local_end, y=y, w=w)\n\n    def _load_batch_from_plan_item(self, item: BatchPlanItem) -> PackedSeqBatch:\n        parts = [self._load_batch_part(part) for part in item.parts]\n        if not parts:\n            raise RuntimeError(\"Empty batch plan item\")\n        if item.pad_size > 0:\n            parts[-1] = self._pad_batch_part(parts[-1], item.pad_size)\n        return PackedSeqBatch(parts=tuple(parts))\n\n    def __iter__(self):\n        plan = self._build_rank_plan()\n        if not plan:\n            return\n        if self.loader_threads <= 1 and self.prefetch_batches <= 1:\n            for item in plan:\n                yield self._load_batch_from_plan_item(item)\n            return\n\n        submit_ahead = max(self.prefetch_batches, self.loader_threads)\n        futures: List[Future] = []\n        next_idx = 0\n        with ThreadPoolExecutor(max_workers=self.loader_threads) as executor:\n            while next_idx < len(plan) and len(futures) < submit_ahead:\n                futures.append(executor.submit(self._load_batch_from_plan_item, plan[next_idx]))\n                next_idx += 1\n            while futures:\n                fut = futures.pop(0)\n                yield fut.result()\n                if next_idx < len(plan):\n                    futures.append(executor.submit(self._load_batch_from_plan_item, plan[next_idx]))\n                    next_idx += 1\n\n\ndef move_batch_to_device(\n    batch: PackedSeqBatch,\n    device: torch.device,\n    copy_non_blocking: bool = False,\n) -> PackedSeqBatch:\n    moved_parts = []\n    for part in batch.parts:\n        moved_parts.append(\n            PackedSeqBatchPart(\n                span_x=part.span_x.to(device, non_blocking=copy_non_blocking),\n                local_end=part.local_end.to(device, non_blocking=copy_non_blocking),\n                y=part.y.to(device, non_blocking=copy_non_blocking),\n                w=part.w.to(device, non_blocking=copy_non_blocking),\n            )\n        )\n    return PackedSeqBatch(parts=tuple(moved_parts))\n\n\nclass DeviceTransferPrefetchLoader:\n    def __init__(\n        self,\n        loader,\n        device: torch.device,\n        prefetch_batches: int = 2,\n        copy_non_blocking: bool = False,\n    ):\n        self.loader = loader\n        self.device = device\n        self.prefetch_batches = max(1, int(prefetch_batches))\n        self.copy_non_blocking = bool(copy_non_blocking)\n\n    def __len__(self) -> int:\n        return len(self.loader)\n\n    def __getattr__(self, name: str):\n        return getattr(self.loader, name)\n\n    def set_epoch(self, epoch: int) -> None:\n        if hasattr(self.loader, \"set_epoch\"):\n            self.loader.set_epoch(epoch)\n\n    def __iter__(self):\n        if self.device.type != \"cuda\":\n            for batch in self.loader:\n                yield batch\n            return\n        result_q: queue.Queue = queue.Queue(maxsize=self.prefetch_batches)\n        sentinel = object()\n        device = self.device\n        stop_event = threading.Event()\n\n        def put_result(item, event) -> bool:\n            while not stop_event.is_set():\n                try:\n                    result_q.put((item, event), timeout=0.1)\n                    return True\n                except queue.Full:\n                    continue\n            return False\n\n        def worker():\n            try:\n                torch.cuda.set_device(device)\n                stream = torch.cuda.Stream(device=device)\n                for batch in self.loader:\n                    if stop_event.is_set():\n                        break\n                    with torch.cuda.stream(stream):\n                        moved = move_batch_to_device(\n                            batch,\n                            device=device,\n                            copy_non_blocking=self.copy_non_blocking,\n                        )\n                        event = torch.cuda.Event()\n                        event.record(stream)\n                    if not put_result(moved, event):\n                        return\n                put_result(sentinel, None)\n            except Exception as exc:  # pragma: no cover - worker thread failure propagation\n                put_result(exc, None)\n\n        thread = threading.Thread(target=worker, daemon=True)\n        thread.start()\n        current_stream = torch.cuda.current_stream(device)\n        try:\n            while True:\n                item, event = result_q.get()\n                if item is sentinel:\n                    break\n                if isinstance(item, Exception):\n                    raise item\n                current_stream.wait_event(event)\n                for part in item.parts:\n                    part.span_x.record_stream(current_stream)\n                    part.local_end.record_stream(current_stream)\n                    part.y.record_stream(current_stream)\n                    part.w.record_stream(current_stream)\n                yield item\n        finally:\n            stop_event.set()\n            thread.join()\n\n\nclass ParallelCNN1D(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int):\n        super().__init__()\n        branch_channels = out_channels // 3\n        self.conv3 = nn.Conv1d(in_channels, branch_channels, kernel_size=3, padding=1)\n        self.conv5 = nn.Conv1d(in_channels, branch_channels, kernel_size=5, padding=2)\n        self.conv7 = nn.Conv1d(in_channels, out_channels - 2 * branch_channels, kernel_size=7, padding=3)\n        self.act = nn.ReLU()\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x_t = x.transpose(1, 2)\n        c3 = self.conv3(x_t)\n        c5 = self.conv5(x_t)\n        c7 = self.conv7(x_t)\n        out = torch.cat([c3, c5, c7], dim=1)\n        return self.act(out).transpose(1, 2)\n\n\nclass FeatureChannelGate(nn.Module):\n    def __init__(self, input_dim: int, hidden_dim: int = 64, init_bias: float = 2.0):\n        super().__init__()\n        hidden_dim = max(1, int(hidden_dim))\n        self.norm = nn.LayerNorm(input_dim)\n        self.fc1 = nn.Linear(input_dim, hidden_dim)\n        self.act = nn.SiLU()\n        self.fc2 = nn.Linear(hidden_dim, input_dim)\n        nn.init.zeros_(self.fc2.weight)\n        nn.init.constant_(self.fc2.bias, float(init_bias))\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        pooled = x.mean(dim=1)\n        gate = self.fc2(self.act(self.fc1(self.norm(pooled))))\n        gate = torch.sigmoid(gate).unsqueeze(1)\n        return x * gate\n\n\nclass GRURegressor(nn.Module):\n    def __init__(\n        self,\n        input_dim: int,\n        hidden_dim: int,\n        num_layers: int,\n        dropout: float,\n        pooling: str = \"last\",\n        bidirectional: bool = False,\n        use_cnn1d: bool = False,\n        input_gate_hidden_dim: int = 0,\n        input_gate_bias: float = 2.0,\n    ):\n        super().__init__()\n        self.pooling = str(pooling).lower()\n        if self.pooling not in {\"last\", \"attn\"}:\n            raise ValueError(f\"Unsupported pooling: {pooling}\")\n        self.bidirectional = bool(bidirectional)\n        self.use_cnn1d = bool(use_cnn1d)\n        self.input_gate_hidden_dim = max(0, int(input_gate_hidden_dim))\n        self.output_dim = int(hidden_dim) * (2 if self.bidirectional else 1)\n        if self.input_gate_hidden_dim > 0:\n            self.input_gate = FeatureChannelGate(\n                input_dim=input_dim,\n                hidden_dim=self.input_gate_hidden_dim,\n                init_bias=float(input_gate_bias),\n            )\n        else:\n            self.input_gate = nn.Identity()\n\n        if self.use_cnn1d:\n            self.cnn = ParallelCNN1D(input_dim, hidden_dim)\n            rnn_input_dim = hidden_dim\n        else:\n            self.cnn = nn.Identity()\n            rnn_input_dim = input_dim\n\n        self.rnn = nn.GRU(\n            input_size=rnn_input_dim,\n            hidden_size=hidden_dim,\n            num_layers=num_layers,\n            dropout=dropout if num_layers > 1 else 0.0,\n            batch_first=True,\n            bidirectional=self.bidirectional,\n        )\n        if self.pooling == \"attn\":\n            self.attn_norm = nn.LayerNorm(self.output_dim)\n            self.attn_proj = nn.Linear(self.output_dim, 1)\n        head_hidden_dim = max(1, self.output_dim // 2)\n        self.head = nn.Sequential(\n            nn.LayerNorm(self.output_dim),\n            nn.Linear(self.output_dim, head_hidden_dim),\n            nn.ReLU(),\n            nn.Linear(head_hidden_dim, 1),\n        )\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x = self.input_gate(x)\n        if self.use_cnn1d:\n            x = self.cnn(x)\n        out, _ = self.rnn(x)\n        if self.pooling == \"attn\":\n            score = self.attn_proj(self.attn_norm(out)).squeeze(-1)\n            weight = torch.softmax(score, dim=1).unsqueeze(-1)\n            pooled = (out * weight).sum(dim=1)\n        else:\n            pooled = out[:, -1, :]\n        return self.head(pooled).squeeze(-1)\n\n\ndef weighted_mse(pred: torch.Tensor, y: torch.Tensor, w: torch.Tensor) -> torch.Tensor:\n    w = torch.clamp(w, min=0.0)\n    w_norm = w / torch.clamp(w.mean(), min=1e-6)\n    return (w_norm * (pred - y) ** 2).mean()\n\n\ndef build_val_epoch_checkpoint_path(out_root: Path, epoch: int) -> Path:\n    return out_root / f\"gru_seq_memmap_val_epoch_{int(epoch):03d}.pt\"\n\n\ndef refresh_top_val_checkpoints(\n    top_val_checkpoints: List[dict],\n    candidate_epoch: int,\n    candidate_val_metrics: dict,\n    keep_topk: int,\n    out_root: Path,\n) -> Tuple[List[dict], bool, List[dict]]:\n    keep_topk = max(1, int(keep_topk))\n    candidate = {\n        \"epoch\": int(candidate_epoch),\n        \"val_ic\": float(candidate_val_metrics[\"ic\"]),\n        \"val_loss\": float(candidate_val_metrics[\"loss\"]),\n        \"path\": str(build_val_epoch_checkpoint_path(out_root, candidate_epoch)),\n    }\n    updated = list(top_val_checkpoints)\n    updated.append(candidate)\n    updated.sort(key=lambda item: (-float(item[\"val_ic\"]), int(item[\"epoch\"])))\n    kept = [dict(item) for item in updated[:keep_topk]]\n    dropped = [dict(item) for item in updated[keep_topk:]]\n    kept_epochs = {int(item[\"epoch\"]) for item in kept}\n    entered_topk = int(candidate_epoch) in kept_epochs\n    for rank_i, item in enumerate(kept, start=1):\n        item[\"rank\"] = int(rank_i)\n    return kept, entered_topk, dropped\n\n\nclass CorrStats:\n    def __init__(self):\n        self.buf: torch.Tensor | None = None\n\n    def update(self, pred: torch.Tensor, y: torch.Tensor, w: torch.Tensor | None = None):\n        p = pred.detach().float().reshape(-1)\n        t = y.detach().float().reshape(-1)\n        mask = torch.isfinite(p) & torch.isfinite(t)\n        if w is not None:\n            ww = w.detach().float().reshape(-1)\n            mask = mask & torch.isfinite(ww) & (ww > 0)\n        zero = torch.zeros_like(p)\n        p = torch.where(mask, p, zero).to(dtype=torch.float64)\n        t = torch.where(mask, t, zero).to(dtype=torch.float64)\n        d = p - t\n        sums = torch.stack(\n            [\n                mask.to(dtype=torch.float64).sum(),\n                p.sum(),\n                t.sum(),\n                (p * p).sum(),\n                (t * t).sum(),\n                (p * t).sum(),\n                torch.abs(d).sum(),\n                (d * d).sum(),\n            ]\n        )\n        if self.buf is None:\n            self.buf = sums\n        else:\n            self.buf = self.buf + sums\n\n    def to_tensor(self, device: torch.device) -> torch.Tensor:\n        if self.buf is None:\n            return torch.zeros((8,), dtype=torch.float64, device=device)\n        return self.buf.to(device=device, dtype=torch.float64)\n\n\ndef metrics_from_tensor(t: torch.Tensor) -> dict:\n    n, sum_p, sum_y, sum_pp, sum_yy, sum_py, sum_abs, sum_sq = [float(x) for x in t.tolist()]\n    if (not np.isfinite(n)) or n <= 0:\n        return {\"rmse\": float(\"nan\"), \"mae\": float(\"nan\"), \"ic\": float(\"nan\"), \"n\": 0}\n    mse = sum_sq / n\n    var_p = max(sum_pp / n - (sum_p / n) ** 2, 1e-12)\n    var_y = max(sum_yy / n - (sum_y / n) ** 2, 1e-12)\n    cov = sum_py / n - (sum_p / n) * (sum_y / n)\n    ic = cov / math.sqrt(var_p * var_y)\n    return {\"rmse\": float(math.sqrt(mse)), \"mae\": float(sum_abs / n), \"ic\": float(ic), \"n\": int(n)}\n\n\ndef compute_feature_stats(\n    train_days: List[Path],\n    prefer_fp16: bool,\n    sample_stride: int = 50,\n    chunk_sample_rows: int = 250_000,\n) -> Tuple[np.ndarray, np.ndarray]:\n    s = None\n    s2 = None\n    n = 0\n    for day in train_days:\n        store = load_day_store(day, prefer_fp16=prefer_fp16)\n        chunk_span = max(int(sample_stride), int(sample_stride) * max(1, int(chunk_sample_rows)))\n        for start in range(0, store.n_rows, chunk_span):\n            stop = min(store.n_rows, start + chunk_span)\n            # Stream smaller memmap slices to reduce pressure and avoid giant one-shot reads.\n            arr = np.array(store.x[start:stop:sample_stride], dtype=np.float32, copy=True)\n            if arr.size == 0:\n                continue\n            arr = np.nan_to_num(arr, nan=0.0, posinf=0.0, neginf=0.0)\n            if s is None:\n                s = arr.sum(axis=0, dtype=np.float64)\n                s2 = (arr * arr).sum(axis=0, dtype=np.float64)\n            else:\n                s += arr.sum(axis=0, dtype=np.float64)\n                s2 += (arr * arr).sum(axis=0, dtype=np.float64)\n            n += arr.shape[0]\n    mean = (s / max(n, 1)).astype(np.float32)\n    var = (s2 / max(n, 1) - mean.astype(np.float64) ** 2).astype(np.float32)\n    std = np.sqrt(np.clip(var, 1e-8, None)).astype(np.float32)\n    return mean, std\n\n\ndef build_feature_stats_cache_path(\n    output_root: Path,\n    row_root: Path,\n    train_days: List[Path],\n    prefer_fp16: bool,\n    sample_stride: int,\n) -> Path:\n    cache_dir = ensure_dir(output_root / \"training_seq\" / \"_feature_stats_cache\")\n    key_src = \"\\n\".join(\n        [\n            str(row_root),\n            str(bool(prefer_fp16)),\n            str(int(sample_stride)),\n            *[d.name for d in train_days],\n        ]\n    )\n    key = hashlib.sha1(key_src.encode(\"utf-8\")).hexdigest()[:16]\n    return cache_dir / f\"feature_stats_{key}.npz\"\n\n\ndef load_feature_stats_cache(path: Path) -> Tuple[np.ndarray, np.ndarray]:\n    with np.load(path) as arr:\n        mean = arr[\"mean\"].astype(np.float32, copy=False)\n        std = arr[\"std\"].astype(np.float32, copy=False)\n    return mean, std\n\n\ndef wait_for_feature_stats_cache(path: Path, timeout_seconds: int = 7200) -> Tuple[np.ndarray, np.ndarray]:\n    deadline = time.time() + max(1, int(timeout_seconds))\n    last_err: Exception | None = None\n    while time.time() < deadline:\n        if path.exists() and path.stat().st_size > 0:\n            try:\n                return load_feature_stats_cache(path)\n            except Exception as exc:  # pragma: no cover - transient partial-write case\n                last_err = exc\n        time.sleep(2.0)\n    if last_err is not None:\n        raise TimeoutError(f\"Timed out waiting for feature stats cache {path}: {last_err}\") from last_err\n    raise TimeoutError(f\"Timed out waiting for feature stats cache {path}\")\n\n\ndef evaluate(\n    model,\n    loader,\n    feat_mean,\n    feat_std,\n    seq_offsets,\n    accelerator,\n    split_name: str = \"eval\",\n    progress_parts: int = 4,\n    copy_non_blocking: bool = False,\n    collect_timing: bool = False,\n) -> dict:\n    model.eval()\n    stats = CorrStats()\n    loss_sum = torch.zeros((), dtype=torch.float64, device=accelerator.device)\n    loss_count = torch.zeros((), dtype=torch.float64, device=accelerator.device)\n    timing_labels = [\n        \"batch_wait_sec\",\n        \"materialize_sec\",\n        \"forward_sec\",\n        \"stats_update_sec\",\n        \"step_total_sec\",\n    ]\n    timing_sums = np.zeros((len(timing_labels),), dtype=np.float64)\n    local_batches = 0\n    total_batches = 0\n    try:\n        total_batches = len(loader)\n    except TypeError:\n        total_batches = 0\n    progress_step = max(1, total_batches // max(1, int(progress_parts))) if total_batches > 0 else 0\n    with torch.no_grad():\n        loader_iter = iter(loader)\n        batch_i = 0\n        while True:\n            t_step0 = time.perf_counter()\n            t0 = time.perf_counter()\n            try:\n                batch = next(loader_iter)\n            except StopIteration:\n                break\n            t1 = time.perf_counter()\n            xb, yb, wb = materialize_batch(\n                batch,\n                seq_offsets=seq_offsets,\n                feat_mean=feat_mean,\n                feat_std=feat_std,\n                device=accelerator.device,\n                copy_non_blocking=copy_non_blocking,\n            )\n            if collect_timing:\n                sync_if_needed(accelerator.device)\n            t2 = time.perf_counter()\n            with accelerator.autocast():\n                pred = model(xb)\n                loss = weighted_mse(pred, yb, wb)\n            if collect_timing:\n                sync_if_needed(accelerator.device)\n            t3 = time.perf_counter()\n            stats.update(pred, yb, wb)\n            if collect_timing:\n                sync_if_needed(accelerator.device)\n            t4 = time.perf_counter()\n            loss_sum = loss_sum + loss.detach().to(dtype=torch.float64)\n            loss_count = loss_count + 1.0\n            local_batches += 1\n            batch_i += 1\n            if collect_timing:\n                timing_sums += np.asarray(\n                    [\n                        t1 - t0,\n                        t2 - t1,\n                        t3 - t2,\n                        t4 - t3,\n                        t4 - t_step0,\n                    ],\n                    dtype=np.float64,\n                )\n            if (\n                accelerator.is_main_process\n                and total_batches > 0\n                and (batch_i % progress_step == 0 or batch_i == total_batches)\n            ):\n                print(f\"[{split_name}] progress {batch_i}/{total_batches}\", flush=True)\n    loss_pair = accelerator.reduce(torch.stack([loss_sum, loss_count]), reduction=\"sum\")\n    loss_den = float(loss_pair[1].item())\n    loss_val = float((loss_pair[0] / loss_den).item()) if loss_den > 0 else float(\"nan\")\n    stats_t = accelerator.reduce(stats.to_tensor(accelerator.device), reduction=\"sum\")\n    m = metrics_from_tensor(stats_t)\n    m[\"loss\"] = float(loss_val)\n    if collect_timing:\n        m[\"_timing\"] = summarize_timing_across_ranks(\n            accelerator=accelerator,\n            labels=timing_labels,\n            count=local_batches,\n            sums=timing_sums,\n            count_key=\"measured_batches\",\n            global_batch=int(loader.batch_size) * int(accelerator.num_processes),\n        )\n    return m\n\n\ndef materialize_batch(\n    batch: PackedSeqBatch,\n    seq_offsets: torch.Tensor,\n    feat_mean: torch.Tensor,\n    feat_std: torch.Tensor,\n    device: torch.device,\n    copy_non_blocking: bool = False,\n) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:\n    xb_parts = []\n    y_parts = []\n    w_parts = []\n    for part in batch.parts:\n        span_x = part.span_x.to(device, non_blocking=copy_non_blocking)\n        local_end = part.local_end.to(device, non_blocking=copy_non_blocking)\n        yb = part.y.to(device, non_blocking=copy_non_blocking)\n        wb = part.w.to(device, non_blocking=copy_non_blocking)\n        seq_idx = local_end[:, None] + seq_offsets[None, :]\n        xb = span_x[seq_idx].float()\n        xb = (xb - feat_mean) / feat_std\n        xb_parts.append(xb)\n        y_parts.append(yb)\n        w_parts.append(wb)\n    if len(xb_parts) == 1:\n        return xb_parts[0], y_parts[0], w_parts[0]\n    return torch.cat(xb_parts, dim=0), torch.cat(y_parts, dim=0), torch.cat(w_parts, dim=0)\n\n\ndef sync_if_needed(device: torch.device) -> None:\n    if device.type == \"cuda\":\n        torch.cuda.synchronize(device)\n\n\ndef warmup_collective(accelerator) -> float:\n    if int(getattr(accelerator, \"num_processes\", 1)) <= 1:\n        return 0.0\n    t0 = time.perf_counter()\n    dummy = torch.zeros((1,), dtype=torch.float32, device=accelerator.device)\n    _ = accelerator.reduce(dummy, reduction=\"sum\")\n    sync_if_needed(accelerator.device)\n    return float(time.perf_counter() - t0)\n\n\ndef summarize_timing_across_ranks(\n    accelerator: Accelerator,\n    labels: List[str],\n    count: int,\n    sums: np.ndarray,\n    count_key: str,\n    global_batch: int | None = None,\n) -> dict:\n    local = torch.tensor([float(count), *sums.tolist()], dtype=torch.float64, device=accelerator.device)[None, :]\n    gathered = accelerator.gather(local)\n    empty = {\n        count_key: int(count),\n        \"per_rank\": [],\n        \"global_step_sec_estimate\": float(\"nan\"),\n        \"global_samples_per_sec_estimate\": float(\"nan\"),\n    }\n    if not accelerator.is_main_process:\n        return empty\n    gathered_np = gathered.detach().cpu().numpy()\n    per_rank = []\n    max_step_sec = 0.0\n    step_label = labels[-1] if labels else None\n    for rank_i, row in enumerate(gathered_np):\n        rank_count = max(1.0, float(row[0]))\n        stats = {\"rank\": int(rank_i), count_key: int(row[0])}\n        for label_i, label in enumerate(labels, start=1):\n            stats[label] = float(row[label_i] / rank_count)\n        if step_label is not None:\n            max_step_sec = max(max_step_sec, float(stats[step_label]))\n        per_rank.append(stats)\n    out = {\n        count_key: int(count),\n        \"per_rank\": per_rank,\n        \"global_step_sec_estimate\": float(max_step_sec) if step_label is not None else float(\"nan\"),\n        \"global_samples_per_sec_estimate\": float(\"nan\"),\n    }\n    if step_label is not None and global_batch is not None:\n        out[\"global_samples_per_sec_estimate\"] = float(global_batch / max(max_step_sec, 1e-9))\n    return out\n\n\ndef profile_train_steps(\n    model,\n    optimizer,\n    loader,\n    accelerator,\n    feat_mean,\n    feat_std,\n    seq_offsets,\n    copy_non_blocking: bool,\n    warmup_steps: int,\n    profile_steps: int,\n    grad_accum_steps: int = 1,\n) -> dict:\n    total_batches = 0\n    try:\n        total_batches = len(loader)\n    except TypeError:\n        total_batches = 0\n    warmup = max(0, int(warmup_steps))\n    active = max(0, int(profile_steps))\n    if total_batches <= 0 or active <= 0:\n        return {\n            \"warmup_steps\": warmup,\n            \"profile_steps\": active,\n            \"measured_steps\": 0,\n            \"per_rank\": [],\n            \"global_step_sec_estimate\": float(\"nan\"),\n            \"global_samples_per_sec_estimate\": float(\"nan\"),\n        }\n    run_steps = min(total_batches, warmup + active)\n    measured_steps = max(0, run_steps - warmup)\n    labels = [\n        \"batch_wait_sec\",\n        \"materialize_sec\",\n        \"forward_sec\",\n        \"backward_sec\",\n        \"optimizer_step_sec\",\n        \"step_total_sec\",\n    ]\n    sums = np.zeros((len(labels),), dtype=np.float64)\n    model.train()\n    train_iter = iter(loader)\n    warmup_collective(accelerator)\n    accum_steps = max(1, int(grad_accum_steps))\n    accum_micro = 0\n    optimizer.zero_grad(set_to_none=True)\n    for step_i in range(1, run_steps + 1):\n        t_step0 = time.perf_counter()\n        t0 = time.perf_counter()\n        batch = next(train_iter)\n        t1 = time.perf_counter()\n        accum_micro += 1\n        last_batch = step_i == run_steps\n        should_step = accum_micro >= accum_steps or last_batch\n        group_total = min(accum_steps, accum_micro + max(0, run_steps - step_i))\n        xb, yb, wb = materialize_batch(\n            batch,\n            seq_offsets=seq_offsets,\n            feat_mean=feat_mean,\n            feat_std=feat_std,\n            device=accelerator.device,\n            copy_non_blocking=copy_non_blocking,\n        )\n        sync_if_needed(accelerator.device)\n        t2 = time.perf_counter()\n        sync_ctx = model.no_sync() if hasattr(model, \"no_sync\") and not should_step else nullcontext()\n        with sync_ctx:\n            with accelerator.autocast():\n                pred = model(xb)\n                loss_raw = weighted_mse(pred, yb, wb)\n                loss = loss_raw / float(group_total)\n            sync_if_needed(accelerator.device)\n            t3 = time.perf_counter()\n            accelerator.backward(loss)\n            sync_if_needed(accelerator.device)\n            t4 = time.perf_counter()\n        if should_step:\n            accelerator.step_optimizer(optimizer)\n            optimizer.zero_grad(set_to_none=True)\n            accum_micro = 0\n        sync_if_needed(accelerator.device)\n        t5 = time.perf_counter()\n        if step_i > warmup:\n            sums += np.asarray(\n                [\n                    t1 - t0,\n                    t2 - t1,\n                    t3 - t2,\n                    t4 - t3,\n                    t5 - t4,\n                    t5 - t_step0,\n                ],\n                dtype=np.float64,\n            )\n        if accelerator.is_main_process and (step_i == run_steps or step_i % max(1, run_steps // 4) == 0):\n            print(f\"[profile-train] step {step_i}/{run_steps} loss={loss_raw.detach().float().item():.6f}\", flush=True)\n    global_batch = int(loader.batch_size) * int(accelerator.num_processes)\n    summary = summarize_timing_across_ranks(\n        accelerator=accelerator,\n        labels=labels,\n        count=measured_steps,\n        sums=sums,\n        count_key=\"measured_steps\",\n        global_batch=global_batch,\n    )\n    summary[\"warmup_steps\"] = warmup\n    summary[\"profile_steps\"] = active\n    return summary\n\n\ndef run(args):\n    run_t0 = time.perf_counter()\n    phase_timings: Dict[str, float] = {}\n    phase_details: Dict[str, dict] = {}\n    data_prepare_t0 = time.perf_counter()\n    set_global_seed(args.seed)\n    cfg = load_config(args.config)\n    train_cfg = dict(cfg[\"train\"])\n    if args.override_epochs > 0:\n        train_cfg[\"epochs\"] = args.override_epochs\n    if args.override_lr > 0:\n        train_cfg[\"learning_rate\"] = args.override_lr\n    if args.override_batch_size > 0:\n        train_cfg[\"batch_size\"] = args.override_batch_size\n\n    row_root = resolve_row_root(cfg, args.cache_name, row_root_override=args.row_root_override)\n    profile_phases = bool(args.profile_phases)\n    profile_warmup_steps = max(0, int(args.profile_warmup_steps))\n    profile_steps = max(0, int(args.profile_steps))\n    stop_after_profile = bool(args.stop_after_profile)\n    profile_only = bool(stop_after_profile and profile_steps > 0)\n    grad_accum_steps = max(1, int(args.grad_accum_steps))\n    split_view_meta = load_split_view_meta(row_root)\n    if split_view_meta is not None:\n        train_days = resolve_named_day_dirs(row_root, split_view_meta.get(\"train_days\"))\n        valid_days = resolve_named_day_dirs(row_root, split_view_meta.get(\"valid_days\"))\n        test_days = resolve_named_day_dirs(row_root, split_view_meta.get(\"test_days\"))\n        train_days = limit_days(train_days, args.train_day_limit)\n        valid_days = limit_days(valid_days, args.valid_day_limit)\n        test_days = limit_days(test_days, args.test_day_limit)\n        train_pool_days = list(train_days) + list(valid_days)\n        days = list(train_days) + list(valid_days) + list(test_days)\n        split_ratio = float(args.valid_split_ratio) if args.valid_split_ratio > 0.0 else float(train_cfg[\"train_split_ratio\"])\n        train_cfg[\"train_split_ratio\"] = split_ratio\n    else:\n        days = list_ready_days(row_root)\n        use_explicit_split = bool(\n            args.train_pool_start_date or args.train_pool_end_date or args.test_start_date or args.test_end_date\n        )\n        if args.max_days > 0 and not use_explicit_split:\n            days = days[: args.max_days]\n        if args.test_start_date or args.test_end_date:\n            test_days = filter_days_by_date(days, args.test_start_date, args.test_end_date)\n        else:\n            test_days = []\n        test_day_names = {d.name for d in test_days}\n        raw_train_pool_days = filter_days_by_date(days, args.train_pool_start_date, args.train_pool_end_date)\n        if args.train_pool_start_date or args.train_pool_end_date:\n            train_pool_days = [d for d in raw_train_pool_days if d.name not in test_day_names]\n        else:\n            train_pool_days = [d for d in days if d.name not in test_day_names]\n        split_ratio = float(args.valid_split_ratio) if args.valid_split_ratio > 0.0 else float(train_cfg[\"train_split_ratio\"])\n        train_cfg[\"train_split_ratio\"] = split_ratio\n        train_days, valid_days = split_days(train_pool_days, split_ratio)\n        train_days = limit_days(train_days, args.train_day_limit)\n        valid_days = limit_days(valid_days, args.valid_day_limit)\n        test_days = limit_days(test_days, args.test_day_limit)\n    if not train_pool_days:\n        raise RuntimeError(\"train_pool_days is empty after applying date filters.\")\n    if len(train_days) == 0 or len(valid_days) == 0:\n        raise RuntimeError(\"Need both train and valid day split.\")\n    print(\n        f\"[data] days_total={len(days)} train_pool_days={len(train_pool_days)} \"\n        f\"train_days={len(train_days)} valid_days={len(valid_days)} test_days={len(test_days)}\",\n        flush=True,\n    )\n    print(f\"[row-root] {row_root}\", flush=True)\n    if args.train_pool_start_date or args.train_pool_end_date:\n        print(\n            f\"[data] train_pool_date_window=[{args.train_pool_start_date or 'min'}, \"\n            f\"{args.train_pool_end_date or 'max'}]\",\n            flush=True,\n        )\n    if args.test_start_date or args.test_end_date:\n        print(\n            f\"[data] test_date_window=[{args.test_start_date or 'min'}, {args.test_end_date or 'max'}]\",\n            flush=True,\n        )\n    data_prepare_local_sec = time.perf_counter() - data_prepare_t0\n    accelerator_init_t0 = time.perf_counter()\n    accelerator = build_runtime(\n        args.runtime_backend,\n        bool(args.use_amp),\n        enable_static_graph=(grad_accum_steps == 1),\n    )\n    phase_timings[\"data_prepare_sec\"] = float(data_prepare_local_sec)\n    phase_timings[\"accelerator_init_sec\"] = float(time.perf_counter() - accelerator_init_t0)\n    loader_threads = max(1, int(args.num_workers))\n    eval_batch_size = int(args.eval_batch_size) if int(args.eval_batch_size) > 0 else int(train_cfg[\"batch_size\"])\n    device_transfer_mode = str(args.device_transfer_mode).strip().lower()\n    device_transfer_prefetch_batches = max(1, int(args.device_transfer_prefetch_batches))\n    # Keep host tensors pinned even on the threaded transfer path so H2D copies\n    # can overlap with compute on the prefetch stream.\n    use_pinned_transfer = bool(train_cfg.get(\"pin_memory\", False))\n    if int(args.loader_prefetch_batches) > 0:\n        loader_prefetch = max(1, int(args.loader_prefetch_batches))\n    else:\n        loader_prefetch = max(1, int(train_cfg.get(\"prefetch_factor\", 2)))\n    train_batch_overlap = max(0, int(args.train_batch_overlap))\n    skip_train_eval = bool(args.skip_train_eval)\n    skip_test = bool(args.skip_test)\n    print(\n        f\"[device-transfer] mode={device_transfer_mode} loader_pin_memory={int(use_pinned_transfer)} \"\n        f\"prefetch_batches={device_transfer_prefetch_batches}\",\n        flush=True,\n    )\n    if train_batch_overlap > 0:\n        print(\n            f\"[train-loader] batch_overlap={train_batch_overlap} batch_stride={int(train_cfg['batch_size']) - train_batch_overlap}\",\n            flush=True,\n        )\n    end_index_cache_dir = ensure_dir(Path(cfg[\"paths\"][\"output_root\"]) / \"training_seq\" / \"_end_index_cache\")\n    loader_init_t0 = time.perf_counter()\n    train_loader = SeqMemmapBatchLoader(\n        train_days,\n        seq_len=args.seq_len,\n        sample_stride=args.sample_stride,\n        batch_size=int(train_cfg[\"batch_size\"]),\n        shuffle=True,\n        max_samples=args.max_samples,\n        prefer_fp16=bool(args.prefer_fp16),\n        rank=accelerator.process_index,\n        world_size=accelerator.num_processes,\n        seed=args.seed,\n        pin_memory=use_pinned_transfer,\n        loader_threads=loader_threads,\n        prefetch_batches=loader_prefetch,\n        pad_last_batch=(accelerator.num_processes > 1),\n        batch_overlap=train_batch_overlap,\n        index_cache_dir=end_index_cache_dir,\n    )\n    phase_timings[\"train_loader_init_sec\"] = float(time.perf_counter() - loader_init_t0)\n    loader_init_t0 = time.perf_counter()\n    train_eval_loader = (\n        SeqMemmapBatchLoader(\n            train_days,\n            seq_len=args.seq_len,\n            sample_stride=args.sample_stride,\n            batch_size=eval_batch_size,\n            shuffle=False,\n            max_samples=args.max_samples,\n            prefer_fp16=bool(args.prefer_fp16),\n            rank=accelerator.process_index,\n            world_size=accelerator.num_processes,\n            seed=args.seed,\n            pin_memory=use_pinned_transfer,\n            loader_threads=loader_threads,\n            prefetch_batches=loader_prefetch,\n            pad_last_batch=False,\n            batch_overlap=0,\n            index_cache_dir=end_index_cache_dir,\n        )\n        if not profile_only and not skip_train_eval and train_batch_overlap > 0\n        else None\n    )\n    phase_timings[\"train_eval_loader_init_sec\"] = float(time.perf_counter() - loader_init_t0)\n    loader_init_t0 = time.perf_counter()\n    valid_loader = (\n        SeqMemmapBatchLoader(\n            valid_days,\n            seq_len=args.seq_len,\n            sample_stride=args.sample_stride,\n            batch_size=eval_batch_size,\n            shuffle=False,\n            max_samples=args.max_samples,\n            prefer_fp16=bool(args.prefer_fp16),\n            rank=accelerator.process_index,\n            world_size=accelerator.num_processes,\n            seed=args.seed,\n            pin_memory=use_pinned_transfer,\n            loader_threads=loader_threads,\n            prefetch_batches=loader_prefetch,\n            pad_last_batch=False,\n            batch_overlap=0,\n            index_cache_dir=end_index_cache_dir,\n        )\n        if not profile_only\n        else None\n    )\n    phase_timings[\"valid_loader_init_sec\"] = float(time.perf_counter() - loader_init_t0)\n    loader_init_t0 = time.perf_counter()\n    test_loader = (\n        SeqMemmapBatchLoader(\n            test_days,\n            seq_len=args.seq_len,\n            sample_stride=args.sample_stride,\n            batch_size=eval_batch_size,\n            shuffle=False,\n            max_samples=-1,\n            prefer_fp16=bool(args.prefer_fp16),\n            rank=accelerator.process_index,\n            world_size=accelerator.num_processes,\n            seed=args.seed,\n            pin_memory=use_pinned_transfer,\n            loader_threads=loader_threads,\n            prefetch_batches=loader_prefetch,\n            pad_last_batch=False,\n            batch_overlap=0,\n            index_cache_dir=end_index_cache_dir,\n        )\n        if test_days and not skip_test and not profile_only\n        else None\n    )\n    if device_transfer_mode == \"thread_prefetch\":\n        train_loader = DeviceTransferPrefetchLoader(\n            train_loader,\n            device=accelerator.device,\n            prefetch_batches=device_transfer_prefetch_batches,\n            copy_non_blocking=use_pinned_transfer,\n        )\n        if train_eval_loader is not None:\n            train_eval_loader = DeviceTransferPrefetchLoader(\n                train_eval_loader,\n                device=accelerator.device,\n                prefetch_batches=device_transfer_prefetch_batches,\n                copy_non_blocking=use_pinned_transfer,\n            )\n        if valid_loader is not None:\n            valid_loader = DeviceTransferPrefetchLoader(\n                valid_loader,\n                device=accelerator.device,\n                prefetch_batches=device_transfer_prefetch_batches,\n                copy_non_blocking=use_pinned_transfer,\n            )\n        if test_loader is not None:\n            test_loader = DeviceTransferPrefetchLoader(\n                test_loader,\n                device=accelerator.device,\n                prefetch_batches=device_transfer_prefetch_batches,\n                copy_non_blocking=use_pinned_transfer,\n            )\n    phase_timings[\"test_loader_init_sec\"] = float(time.perf_counter() - loader_init_t0)\n    stats_stride = max(20, args.sample_stride)\n    feature_stats_cache = build_feature_stats_cache_path(\n        output_root=Path(cfg[\"paths\"][\"output_root\"]),\n        row_root=row_root,\n        train_days=train_days,\n        prefer_fp16=bool(args.prefer_fp16),\n        sample_stride=stats_stride,\n    )\n    feature_stats_t0 = time.perf_counter()\n    if accelerator.is_main_process:\n        if feature_stats_cache.exists() and feature_stats_cache.stat().st_size > 0:\n            print(f\"[feature-stats] reused_from={feature_stats_cache}\", flush=True)\n            feat_mean_np, feat_std_np = load_feature_stats_cache(feature_stats_cache)\n        else:\n            print(\n                f\"[feature-stats] computing cache={feature_stats_cache} sample_stride={stats_stride}\",\n                flush=True,\n            )\n            feat_mean_np, feat_std_np = compute_feature_stats(\n                train_days, prefer_fp16=bool(args.prefer_fp16), sample_stride=stats_stride\n            )\n            tmp_path = feature_stats_cache.with_suffix(f\".tmp.{int(time.time())}.npz\")\n            np.savez(tmp_path, mean=feat_mean_np, std=feat_std_np)\n            tmp_path.replace(feature_stats_cache)\n            print(f\"[feature-stats] saved_cache={feature_stats_cache}\", flush=True)\n    else:\n        print(f\"[feature-stats] waiting_for_cache={feature_stats_cache}\", flush=True)\n        feat_mean_np, feat_std_np = wait_for_feature_stats_cache(feature_stats_cache)\n        print(f\"[feature-stats] loaded_cache={feature_stats_cache}\", flush=True)\n    accelerator.wait_for_everyone()\n    phase_timings[\"feature_stats_sec\"] = float(time.perf_counter() - feature_stats_t0)\n\n    model_init_t0 = time.perf_counter()\n    model = GRURegressor(\n        input_dim=200,\n        hidden_dim=args.hidden_dim,\n        num_layers=args.num_layers,\n        dropout=args.dropout,\n        pooling=args.pooling,\n        bidirectional=bool(args.bidirectional),\n        use_cnn1d=bool(getattr(args, \"use_cnn1d\", 0)),\n        input_gate_hidden_dim=int(getattr(args, \"input_gate_hidden_dim\", 0)),\n        input_gate_bias=float(getattr(args, \"input_gate_bias\", 2.0)),\n    )\n    optimizer = torch.optim.AdamW(\n        model.parameters(), lr=float(train_cfg[\"learning_rate\"]), weight_decay=float(args.weight_decay)\n    )\n    phase_timings[\"model_optimizer_init_sec\"] = float(time.perf_counter() - model_init_t0)\n\n    prepare_t0 = time.perf_counter()\n    if test_loader is not None:\n        model, optimizer = accelerator.prepare(model, optimizer)\n    else:\n        model, optimizer = accelerator.prepare(model, optimizer)\n    phase_timings[\"accelerator_prepare_sec\"] = float(time.perf_counter() - prepare_t0)\n    tensor_init_t0 = time.perf_counter()\n    feat_mean = torch.from_numpy(feat_mean_np).to(accelerator.device)[None, None, :]\n    feat_std = torch.from_numpy(feat_std_np).to(accelerator.device)[None, None, :]\n    seq_offsets = torch.arange(-(args.seq_len - 1), 1, dtype=torch.long, device=accelerator.device)\n    phase_timings[\"device_tensor_init_sec\"] = float(time.perf_counter() - tensor_init_t0)\n\n    output_init_t0 = time.perf_counter()\n    out_root = ensure_dir(Path(cfg[\"paths\"][\"output_root\"]) / \"training_seq\" / args.run_name)\n    log_path = out_root / \"train_log.csv\"\n    best_model_path = out_root / \"gru_seq_memmap_best_val_ic.pt\"\n    save_topk_val_checkpoints = max(\n        1,\n        int(getattr(args, \"save_topk_val_checkpoints\", 1)),\n        int(getattr(args, \"test_checkpoint_rank\", 1)),\n    )\n    selected_test_checkpoint_rank = max(1, int(getattr(args, \"test_checkpoint_rank\", 1)))\n    if accelerator.is_main_process and not profile_only:\n        with log_path.open(\"w\", newline=\"\", encoding=\"utf-8\") as f:\n            csv.writer(f).writerow(\n                [\n                    \"epoch\",\n                    \"lr\",\n                    \"train_loss\",\n                    \"train_ic\",\n                    \"train_rmse\",\n                    \"val_loss\",\n                    \"val_ic\",\n                    \"val_rmse\",\n                    \"val_mae\",\n                    \"is_best\",\n                    \"pooling\",\n                ]\n            )\n    phase_timings[\"output_init_sec\"] = float(time.perf_counter() - output_init_t0)\n\n    epochs = int(train_cfg[\"epochs\"])\n    epoch_lrs = build_epoch_lr_schedule(train_cfg, epochs=epochs)\n    total_train_batches = 0\n    try:\n        total_train_batches = len(train_loader)\n    except TypeError:\n        total_train_batches = 0\n    train_progress_step = (\n        max(1, total_train_batches // max(1, int(args.train_progress_splits))) if total_train_batches > 0 else 0\n    )\n    best_epoch = 0\n    best_val_m: Dict[str, float] | None = None\n    final_val_m: Dict[str, float] | None = None\n    top_val_checkpoints: List[dict] = []\n    first_train_step_local: np.ndarray | None = None\n    collective_warmup_sec = 0.0\n    collective_warmup_done = False\n    train_eval_timing: dict | None = None\n    valid_eval_timing: dict | None = None\n    test_eval_timing: dict | None = None\n    if profile_only:\n        train_loader.set_epoch(1)\n        profile_summary = profile_train_steps(\n            model=model,\n            optimizer=optimizer,\n            loader=train_loader,\n            accelerator=accelerator,\n            feat_mean=feat_mean,\n            feat_std=feat_std,\n            seq_offsets=seq_offsets,\n            copy_non_blocking=use_pinned_transfer,\n            warmup_steps=profile_warmup_steps,\n            profile_steps=profile_steps,\n            grad_accum_steps=grad_accum_steps,\n        )\n        accelerator.wait_for_everyone()\n        if accelerator.is_main_process:\n            summary_path = out_root / \"step_profile_summary.json\"\n            summary = {\n                \"run_name\": args.run_name,\n                \"row_root\": str(row_root),\n                \"world_size\": int(accelerator.num_processes),\n                \"batch_size_per_rank\": int(train_cfg[\"batch_size\"]),\n                \"train_batch_overlap\": train_batch_overlap,\n                \"runtime_backend\": args.runtime_backend,\n                \"device_transfer_mode\": device_transfer_mode,\n                \"seq_len\": int(args.seq_len),\n                \"sample_stride\": int(args.sample_stride),\n                \"profile\": profile_summary,\n            }\n            summary_path.write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding=\"utf-8\")\n            print(\n                f\"[profile-train] saved={summary_path} global_step_sec={profile_summary['global_step_sec_estimate']:.6f} \"\n                f\"global_samples_per_sec={profile_summary['global_samples_per_sec_estimate']:.2f}\",\n                flush=True,\n            )\n        accelerator.close()\n        return\n    train_loop_t0 = time.perf_counter()\n    for epoch in range(1, epochs + 1):\n        current_lr = float(epoch_lrs[epoch - 1])\n        set_optimizer_lr(optimizer, current_lr)\n        epoch_train_t0 = time.perf_counter()\n        model.train()\n        train_loader.set_epoch(epoch)\n        if accelerator.is_main_process:\n            print(f\"[epoch {epoch}] lr={current_lr:.8f}\", flush=True)\n        train_iter = iter(train_loader)\n        if not collective_warmup_done:\n            collective_warmup_sec = warmup_collective(accelerator)\n            collective_warmup_done = True\n            phase_timings[\"collective_warmup_sec\"] = float(collective_warmup_sec)\n        batch_i = 0\n        accum_micro = 0\n        optimizer.zero_grad(set_to_none=True)\n        while True:\n            fetch_t0 = time.perf_counter()\n            try:\n                batch = next(train_iter)\n            except StopIteration:\n                break\n            fetch_t1 = time.perf_counter()\n            batch_i += 1\n            accum_micro += 1\n            last_batch = total_train_batches > 0 and batch_i == total_train_batches\n            should_step = accum_micro >= grad_accum_steps or last_batch\n            group_total = min(grad_accum_steps, accum_micro + max(0, total_train_batches - batch_i))\n            measure_first_step = bool(profile_phases and first_train_step_local is None and epoch == 1)\n            xb, yb, wb = materialize_batch(\n                batch,\n                seq_offsets=seq_offsets,\n                feat_mean=feat_mean,\n                feat_std=feat_std,\n                device=accelerator.device,\n                copy_non_blocking=use_pinned_transfer,\n            )\n            if measure_first_step:\n                sync_if_needed(accelerator.device)\n                mat_t = time.perf_counter()\n            sync_ctx = model.no_sync() if hasattr(model, \"no_sync\") and not should_step else nullcontext()\n            with sync_ctx:\n                with accelerator.autocast():\n                    pred = model(xb)\n                    loss_raw = weighted_mse(pred, yb, wb)\n                    loss = loss_raw / float(group_total)\n                if measure_first_step:\n                    sync_if_needed(accelerator.device)\n                    fwd_t = time.perf_counter()\n                accelerator.backward(loss)\n                if measure_first_step:\n                    sync_if_needed(accelerator.device)\n                    bwd_t = time.perf_counter()\n            if should_step:\n                accelerator.step_optimizer(optimizer)\n                optimizer.zero_grad(set_to_none=True)\n                accum_micro = 0\n            if measure_first_step:\n                sync_if_needed(accelerator.device)\n                step_t = time.perf_counter()\n                first_train_step_local = np.asarray(\n                    [\n                        fetch_t1 - fetch_t0,\n                        mat_t - fetch_t1,\n                        fwd_t - mat_t,\n                        bwd_t - fwd_t,\n                        step_t - bwd_t,\n                        step_t - fetch_t0,\n                    ],\n                    dtype=np.float64,\n                )\n            if (\n                accelerator.is_main_process\n                and total_train_batches > 0\n                and (batch_i % train_progress_step == 0 or batch_i == total_train_batches)\n            ):\n                print(\n                    f\"[epoch {epoch}] train progress {batch_i}/{total_train_batches} \"\n                    f\"loss={loss_raw.detach().float().item():.6f}\",\n                    flush=True,\n                )\n        phase_timings[f\"epoch_{epoch}_train_sec\"] = float(time.perf_counter() - epoch_train_t0)\n\n        if skip_train_eval:\n            train_m = {\n                \"loss\": float(\"nan\"),\n                \"ic\": float(\"nan\"),\n                \"rmse\": float(\"nan\"),\n                \"mae\": float(\"nan\"),\n                \"n\": 0,\n            }\n            if accelerator.is_main_process:\n                print(f\"[epoch {epoch}] train-eval skipped\", flush=True)\n        else:\n            train_eval_t0 = time.perf_counter()\n            train_m = evaluate(\n                model,\n                train_eval_loader if train_eval_loader is not None else train_loader,\n                feat_mean,\n                feat_std,\n                seq_offsets,\n                accelerator,\n                split_name=f\"epoch {epoch} train-eval\",\n                progress_parts=args.eval_progress_splits,\n                copy_non_blocking=use_pinned_transfer,\n                collect_timing=profile_phases,\n            )\n            phase_timings[f\"epoch_{epoch}_train_eval_sec\"] = float(time.perf_counter() - train_eval_t0)\n            train_eval_timing = train_m.pop(\"_timing\", None)\n        valid_eval_t0 = time.perf_counter()\n        val_m = evaluate(\n            model,\n            valid_loader,\n            feat_mean,\n            feat_std,\n            seq_offsets,\n            accelerator,\n            split_name=f\"epoch {epoch} valid\",\n            progress_parts=args.eval_progress_splits,\n            copy_non_blocking=use_pinned_transfer,\n            collect_timing=profile_phases,\n        )\n        phase_timings[f\"epoch_{epoch}_valid_eval_sec\"] = float(time.perf_counter() - valid_eval_t0)\n        valid_eval_timing = val_m.pop(\"_timing\", None)\n        final_val_m = dict(val_m)\n        top_val_checkpoints, entered_topk, dropped_topk = refresh_top_val_checkpoints(\n            top_val_checkpoints=top_val_checkpoints,\n            candidate_epoch=epoch,\n            candidate_val_metrics=val_m,\n            keep_topk=save_topk_val_checkpoints,\n            out_root=out_root,\n        )\n        is_best = 0\n        if best_val_m is None or float(val_m[\"ic\"]) > float(best_val_m[\"ic\"]):\n            best_val_m = dict(val_m)\n            best_epoch = epoch\n            is_best = 1\n            if accelerator.is_main_process:\n                torch.save(accelerator.unwrap_model(model).state_dict(), best_model_path)\n\n        if accelerator.is_main_process:\n            if entered_topk:\n                torch.save(\n                    accelerator.unwrap_model(model).state_dict(),\n                    build_val_epoch_checkpoint_path(out_root, epoch),\n                )\n            for dropped_item in dropped_topk:\n                dropped_path = Path(str(dropped_item[\"path\"]))\n                if dropped_path.exists():\n                    dropped_path.unlink()\n            with log_path.open(\"a\", newline=\"\", encoding=\"utf-8\") as f:\n                csv.writer(f).writerow(\n                    [\n                        epoch,\n                        current_lr,\n                        train_m[\"loss\"],\n                        train_m[\"ic\"],\n                        train_m[\"rmse\"],\n                        val_m[\"loss\"],\n                        val_m[\"ic\"],\n                        val_m[\"rmse\"],\n                        val_m[\"mae\"],\n                        is_best,\n                        str(args.pooling),\n                    ]\n                )\n            print(\n                f\"[epoch {epoch}] \"\n                f\"lr={current_lr:.8f} \"\n                f\"train_loss={train_m['loss']:.6f} train_ic={train_m['ic']:.6f} train_rmse={train_m['rmse']:.6f} \"\n                f\"val_loss={val_m['loss']:.6f} val_ic={val_m['ic']:.6f} val_rmse={val_m['rmse']:.6f} \"\n                f\"best_epoch={best_epoch}\"\n            )\n    phase_timings[\"fit_loop_sec\"] = float(time.perf_counter() - train_loop_t0)\n    if profile_phases and first_train_step_local is not None:\n        phase_details[\"first_train_step\"] = summarize_timing_across_ranks(\n            accelerator=accelerator,\n            labels=[\n                \"batch_wait_sec\",\n                \"materialize_sec\",\n                \"forward_sec\",\n                \"backward_sec\",\n                \"optimizer_step_sec\",\n                \"step_total_sec\",\n            ],\n            count=1,\n            sums=first_train_step_local,\n            count_key=\"measured_steps\",\n            global_batch=int(train_cfg[\"batch_size\"]) * int(accelerator.num_processes),\n        )\n    if collective_warmup_done:\n        phase_details[\"collective_warmup_sec\"] = {\"local_sec\": float(collective_warmup_sec)}\n    if train_eval_timing is not None:\n        phase_details[\"train_eval_timing\"] = train_eval_timing\n    if valid_eval_timing is not None:\n        phase_details[\"valid_eval_timing\"] = valid_eval_timing\n\n    def evaluate_saved_checkpoint(checkpoint_entry: dict, split_name: str) -> Tuple[Dict[str, float], float, dict | None]:\n        eval_t0 = time.perf_counter()\n        state_dict = torch.load(str(checkpoint_entry[\"path\"]), map_location=\"cpu\", weights_only=True)\n        accelerator.unwrap_model(model).load_state_dict(state_dict)\n        accelerator.wait_for_everyone()\n        if accelerator.is_main_process:\n            print(\n                f\"[{split_name}] evaluating val_rank={checkpoint_entry['rank']} \"\n                f\"epoch={checkpoint_entry['epoch']} val_ic={checkpoint_entry['val_ic']:.6f}\",\n                flush=True,\n            )\n        metrics = evaluate(\n            model,\n            test_loader,\n            feat_mean,\n            feat_std,\n            seq_offsets,\n            accelerator,\n            split_name=split_name,\n            progress_parts=args.eval_progress_splits,\n            copy_non_blocking=use_pinned_transfer,\n            collect_timing=profile_phases,\n        )\n        elapsed = float(time.perf_counter() - eval_t0)\n        timing = metrics.pop(\"_timing\", None)\n        return metrics, elapsed, timing\n\n    accelerator.wait_for_everyone()\n    test_m_best: Dict[str, float] | None = None\n    test_m_selected: Dict[str, float] | None = None\n    selected_test_checkpoint: dict | None = None\n    if test_loader is not None and not skip_test:\n        if not top_val_checkpoints:\n            raise RuntimeError(\"No validation checkpoints were recorded for test evaluation.\")\n        if selected_test_checkpoint_rank > len(top_val_checkpoints):\n            raise RuntimeError(\n                f\"Requested test_checkpoint_rank={selected_test_checkpoint_rank} but only \"\n                f\"{len(top_val_checkpoints)} validation checkpoints are available.\"\n            )\n        best_test_checkpoint = top_val_checkpoints[0]\n        selected_test_checkpoint = top_val_checkpoints[selected_test_checkpoint_rank - 1]\n        test_m_best, test_best_sec, test_eval_timing = evaluate_saved_checkpoint(\n            best_test_checkpoint,\n            split_name=\"test best-val\",\n        )\n        phase_timings[\"test_eval_best_val_sec\"] = float(test_best_sec)\n        if selected_test_checkpoint_rank == 1:\n            test_m_selected = dict(test_m_best)\n            phase_timings[\"test_eval_selected_val_rank_sec\"] = float(test_best_sec)\n            phase_timings[\"test_eval_sec\"] = float(test_best_sec)\n        else:\n            test_m_selected, test_selected_sec, test_eval_timing_selected = evaluate_saved_checkpoint(\n                selected_test_checkpoint,\n                split_name=f\"test val-rank-{selected_test_checkpoint_rank}\",\n            )\n            phase_timings[\"test_eval_selected_val_rank_sec\"] = float(test_selected_sec)\n            phase_timings[\"test_eval_sec\"] = float(test_best_sec + test_selected_sec)\n            if test_eval_timing_selected is not None:\n                phase_details[\"test_eval_timing_selected_val_rank\"] = test_eval_timing_selected\n        if accelerator.is_main_process:\n            print(\n                f\"[test best-val] rmse={test_m_best['rmse']:.6f} mae={test_m_best['mae']:.6f} ic={test_m_best['ic']:.6f}\",\n                flush=True,\n            )\n            if selected_test_checkpoint_rank != 1 and test_m_selected is not None:\n                print(\n                    f\"[test val-rank-{selected_test_checkpoint_rank}] \"\n                    f\"rmse={test_m_selected['rmse']:.6f} mae={test_m_selected['mae']:.6f} ic={test_m_selected['ic']:.6f}\",\n                    flush=True,\n                )\n    elif test_loader is not None and accelerator.is_main_process:\n        print(\"[test] skipped\", flush=True)\n    if test_eval_timing is not None:\n        phase_details[\"test_eval_timing\"] = test_eval_timing\n    final_save_t0 = time.perf_counter()\n    if accelerator.is_main_process:\n        torch.save(accelerator.unwrap_model(model).state_dict(), out_root / \"gru_seq_memmap_ddp.pt\")\n    accelerator.wait_for_everyone()\n    phase_timings[\"final_save_sec\"] = float(time.perf_counter() - final_save_t0)\n    phase_timings[\"total_run_sec\"] = float(time.perf_counter() - run_t0)\n    if accelerator.is_main_process:\n        summary = {\n            \"run_name\": args.run_name,\n            \"cache_name\": args.cache_name,\n            \"max_days\": args.max_days,\n            \"train_pool_start_date\": args.train_pool_start_date,\n            \"train_pool_end_date\": args.train_pool_end_date,\n            \"test_start_date\": args.test_start_date,\n            \"test_end_date\": args.test_end_date,\n            \"valid_split_ratio\": split_ratio,\n            \"selection_metric\": \"val_ic\",\n            \"train_day_count\": len(train_days),\n            \"valid_day_count\": len(valid_days),\n            \"test_day_count\": len(test_days),\n            \"train_samples\": int(train_loader.total_samples),\n            \"valid_samples\": int(valid_loader.total_samples),\n            \"test_samples\": int(test_loader.total_samples) if test_loader is not None else 0,\n            \"seed\": int(args.seed),\n            \"seq_len\": args.seq_len,\n            \"sample_stride\": args.sample_stride,\n            \"hidden_dim\": args.hidden_dim,\n            \"num_layers\": args.num_layers,\n            \"bidirectional\": bool(args.bidirectional),\n            \"use_cnn1d\": bool(getattr(args, \"use_cnn1d\", 0)),\n            \"input_gate_hidden_dim\": int(getattr(args, \"input_gate_hidden_dim\", 0)),\n            \"input_gate_bias\": float(getattr(args, \"input_gate_bias\", 2.0)),\n            \"weight_decay\": args.weight_decay,\n            \"prefer_fp16\": bool(args.prefer_fp16),\n            \"use_amp\": bool(args.use_amp),\n            \"pooling\": str(args.pooling),\n            \"loader_backend\": \"threaded_seq_batch_loader\",\n            \"loader_threads\": loader_threads,\n            \"loader_prefetch_batches\": loader_prefetch,\n            \"train_batch_overlap\": train_batch_overlap,\n            \"train_batch_stride\": int(train_cfg[\"batch_size\"]) - train_batch_overlap,\n            \"eval_batch_size_per_rank\": eval_batch_size,\n            \"grad_accum_steps\": grad_accum_steps,\n            \"effective_global_batch\": int(train_cfg[\"batch_size\"]) * int(accelerator.num_processes) * grad_accum_steps,\n            \"runtime_backend\": args.runtime_backend,\n            \"device_transfer_mode\": device_transfer_mode,\n            \"device_transfer_prefetch_batches\": device_transfer_prefetch_batches,\n            \"pinned_nonblocking_transfer\": use_pinned_transfer,\n            \"skip_train_eval\": skip_train_eval,\n            \"skip_test\": skip_test,\n            \"train_cfg\": train_cfg,\n            \"epoch_learning_rates\": [float(x) for x in epoch_lrs],\n            \"loss\": \"weighted_mse\",\n            \"metric_main\": \"val_ic\",\n            \"test_selection_rule\": f\"val_ic_rank_{selected_test_checkpoint_rank}\",\n            \"save_topk_val_checkpoints\": int(save_topk_val_checkpoints),\n            \"best_epoch_by_val_ic\": best_epoch,\n            \"best_val_metrics\": best_val_m,\n            \"final_epoch_val_metrics\": final_val_m,\n            \"top_val_checkpoints\": top_val_checkpoints,\n            \"selected_test_checkpoint_rank\": int(selected_test_checkpoint_rank),\n            \"selected_test_checkpoint_epoch\": int(selected_test_checkpoint[\"epoch\"]) if selected_test_checkpoint is not None else None,\n            \"selected_test_model_path\": str(selected_test_checkpoint[\"path\"]) if selected_test_checkpoint is not None else None,\n            \"test_metrics_at_best_val\": test_m_best,\n            \"test_metrics_at_selected_val_rank\": test_m_selected,\n            \"best_model_path\": str(best_model_path),\n            \"phase_timings\": phase_timings,\n            \"phase_details\": phase_details,\n        }\n        (out_root / \"training_summary.json\").write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding=\"utf-8\")\n        np.savez(out_root / \"feature_stats.npz\", mean=feat_mean_np, std=feat_std_np)\n        if profile_phases:\n            phase_summary = {\n                \"run_name\": args.run_name,\n                \"row_root\": str(row_root),\n                \"world_size\": int(accelerator.num_processes),\n                \"batch_size_per_rank\": int(train_cfg[\"batch_size\"]),\n                \"train_batch_overlap\": train_batch_overlap,\n                \"runtime_backend\": args.runtime_backend,\n                \"device_transfer_mode\": device_transfer_mode,\n                \"phase_timings\": phase_timings,\n                \"phase_details\": phase_details,\n            }\n            phase_path = out_root / \"phase_profile_summary.json\"\n            phase_path.write_text(json.dumps(phase_summary, ensure_ascii=False, indent=2), encoding=\"utf-8\")\n            print(f\"[phase-profile] saved={phase_path}\", flush=True)\n    accelerator.close()\n\n\ndef main():\n    p = argparse.ArgumentParser()\n    p.add_argument(\"--config\", default=\"/intern9/huhongkai/hs300_factor_lab/configs/experiment_2024_2025_memmap.json\")\n    p.add_argument(\"--run-name\", default=\"seq_gru_memmap_run\")\n    p.add_argument(\"--cache-name\", default=\"top200_eps20_rows_2024_2025\")\n    p.add_argument(\n        \"--row-root-override\",\n        type=str,\n        default=\"\",\n        help=\"Optional row_memmap root or exact cache dir. Useful for local staged cache under /dev/shm.\",\n    )\n    p.add_argument(\"--max-days\", type=int, default=-1)\n    p.add_argument(\"--max-samples\", type=int, default=-1)\n    p.add_argument(\"--train-day-limit\", type=int, default=-1)\n    p.add_argument(\"--valid-day-limit\", type=int, default=-1)\n    p.add_argument(\"--test-day-limit\", type=int, default=-1)\n    p.add_argument(\"--train-pool-start-date\", type=str, default=\"\")\n    p.add_argument(\"--train-pool-end-date\", type=str, default=\"\")\n    p.add_argument(\"--test-start-date\", type=str, default=\"\")\n    p.add_argument(\"--test-end-date\", type=str, default=\"\")\n    p.add_argument(\"--valid-split-ratio\", type=float, default=-1.0)\n    p.add_argument(\"--seed\", type=int, default=20260312)\n    p.add_argument(\"--seq-len\", type=int, default=60)\n    p.add_argument(\"--sample-stride\", type=int, default=10)\n    p.add_argument(\"--train-batch-overlap\", type=int, default=0)\n    p.add_argument(\"--eval-batch-size\", type=int, default=0)\n    p.add_argument(\"--grad-accum-steps\", type=int, default=1)\n    p.add_argument(\"--runtime-backend\", type=str, default=\"native\", choices=[\"accelerate\", \"native\"])\n    p.add_argument(\n        \"--device-transfer-mode\",\n        type=str,\n        default=\"thread_prefetch\",\n        choices=[\"direct\", \"thread_prefetch\"],\n    )\n    p.add_argument(\"--device-transfer-prefetch-batches\", type=int, default=2)\n    p.add_argument(\"--loader-prefetch-batches\", type=int, default=0)\n    p.add_argument(\"--prefer-fp16\", type=int, default=1)\n    p.add_argument(\"--override-epochs\", type=int, default=-1)\n    p.add_argument(\"--override-lr\", type=float, default=-1.0)\n    p.add_argument(\"--override-batch-size\", type=int, default=-1)\n    p.add_argument(\"--hidden-dim\", type=int, default=256)\n    p.add_argument(\"--num-layers\", type=int, default=2)\n    p.add_argument(\"--dropout\", type=float, default=0.1)\n    p.add_argument(\"--pooling\", type=str, default=\"last\", choices=[\"last\", \"attn\"])\n    p.add_argument(\"--bidirectional\", type=int, default=0, help=\"1 to enable bidirectional GRU\")\n    p.add_argument(\"--use-cnn1d\", type=int, default=0, help=\"1 to use parallel 1D-CNN before GRU\")\n    p.add_argument(\"--input-gate-hidden-dim\", type=int, default=0, help=\">0 to enable feature/channel gate before GRU\")\n    p.add_argument(\"--input-gate-bias\", type=float, default=2.0, help=\"Initial bias for feature/channel gate\")\n    p.add_argument(\"--weight-decay\", type=float, default=1e-5)\n    p.add_argument(\"--num-workers\", type=int, default=4)\n    p.add_argument(\"--use-amp\", type=int, default=1, help=\"1 to enable fp16 mixed precision\")\n    p.add_argument(\"--train-progress-splits\", type=int, default=10)\n    p.add_argument(\"--eval-progress-splits\", type=int, default=4)\n    p.add_argument(\"--skip-train-eval\", type=int, default=1, help=\"1 to skip full train-set evaluation\")\n    p.add_argument(\"--skip-test\", type=int, default=0, help=\"1 to skip final test evaluation\")\n    p.add_argument(\"--profile-phases\", type=int, default=0)\n    p.add_argument(\"--profile-warmup-steps\", type=int, default=0)\n    p.add_argument(\"--profile-steps\", type=int, default=0)\n    p.add_argument(\"--stop-after-profile\", type=int, default=0)\n    p.add_argument(\"--save-topk-val-checkpoints\", type=int, default=1)\n    p.add_argument(\"--test-checkpoint-rank\", type=int, default=1)\n    args = p.parse_args()\n    run(args)\n\n\nif __name__ == \"__main__\":\n    main()\n","afterFullFileContent":"import argparse\nimport csv\nimport hashlib\nimport json\nimport math\nimport os\nimport queue\nimport random\nimport threading\nimport time\nfrom concurrent.futures import Future, ThreadPoolExecutor\nfrom contextlib import nullcontext\nfrom dataclasses import dataclass\nfrom datetime import timedelta\nfrom pathlib import Path\nfrom typing import Dict, List, Tuple\n\nimport numpy as np\nimport torch\nimport torch.distributed as dist\nimport torch.nn as nn\nfrom torch.nn.parallel import DistributedDataParallel as DDP\nfrom accelerate import Accelerator\n\nfrom common import ensure_dir, load_config, load_split_view_meta, resolve_named_day_dirs, resolve_row_root\n\n\ndef set_global_seed(seed: int) -> None:\n    seed = int(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n\n\ndef build_epoch_lr_schedule(train_cfg: dict, epochs: int) -> List[float]:\n    base_lr = float(train_cfg[\"learning_rate\"])\n    raw_values = train_cfg.get(\"lr_epoch_values\")\n    if raw_values is None:\n        return [base_lr for _ in range(int(epochs))]\n    values = [float(x) for x in raw_values]\n    if len(values) != int(epochs):\n        raise ValueError(\n            f\"train.lr_epoch_values length ({len(values)}) must match epochs ({int(epochs)}).\"\n        )\n    return values\n\n\ndef set_optimizer_lr(optimizer, lr: float) -> None:\n    lr = float(lr)\n    for group in optimizer.param_groups:\n        group[\"lr\"] = lr\n\n\nclass AccelerateRuntime:\n    def __init__(self, use_amp: bool):\n        self.accelerator = Accelerator(mixed_precision=\"fp16\" if bool(use_amp) else \"no\")\n        self.device = self.accelerator.device\n        self.process_index = int(self.accelerator.process_index)\n        self.num_processes = int(self.accelerator.num_processes)\n        self.is_main_process = bool(self.accelerator.is_main_process)\n\n    def autocast(self):\n        return self.accelerator.autocast()\n\n    def backward(self, loss: torch.Tensor) -> None:\n        self.accelerator.backward(loss)\n\n    def step_optimizer(self, optimizer) -> None:\n        optimizer.step()\n\n    def prepare(self, model, optimizer):\n        return self.accelerator.prepare(model, optimizer)\n\n    def unwrap_model(self, model):\n        return self.accelerator.unwrap_model(model)\n\n    def reduce(self, tensor: torch.Tensor, reduction: str = \"sum\") -> torch.Tensor:\n        return self.accelerator.reduce(tensor, reduction=reduction)\n\n    def gather(self, tensor: torch.Tensor) -> torch.Tensor:\n        return self.accelerator.gather(tensor)\n\n    def wait_for_everyone(self) -> None:\n        self.accelerator.wait_for_everyone()\n\n    def close(self) -> None:\n        return None\n\n\nclass NativeDDPRuntime:\n    def __init__(self, use_amp: bool, enable_static_graph: bool = True):\n        if torch.cuda.is_available():\n            local_rank = int(os.environ.get(\"LOCAL_RANK\", 0))\n            torch.cuda.set_device(local_rank)\n            self.device = torch.device(\"cuda\", local_rank)\n        else:\n            self.device = torch.device(\"cpu\")\n        requested_world = int(os.environ.get(\"WORLD_SIZE\", \"1\"))\n        self._distributed = requested_world > 1\n        if self._distributed and not dist.is_initialized():\n            backend = \"nccl\" if self.device.type == \"cuda\" else \"gloo\"\n            dist.init_process_group(backend=backend, timeout=timedelta(seconds=7200))\n        if dist.is_initialized():\n            self.process_index = int(dist.get_rank())\n            self.num_processes = int(dist.get_world_size())\n        else:\n            self.process_index = 0\n            self.num_processes = 1\n        self.is_main_process = self.process_index == 0\n        self.use_amp = bool(use_amp) and self.device.type == \"cuda\"\n        self.enable_static_graph = bool(enable_static_graph)\n        self.scaler = torch.cuda.amp.GradScaler(enabled=self.use_amp)\n\n    def autocast(self):\n        if self.device.type == \"cuda\":\n            return torch.autocast(device_type=\"cuda\", dtype=torch.float16, enabled=self.use_amp)\n        return nullcontext()\n\n    def backward(self, loss: torch.Tensor) -> None:\n        if self.use_amp:\n            self.scaler.scale(loss).backward()\n        else:\n            loss.backward()\n\n    def step_optimizer(self, optimizer) -> None:\n        if self.use_amp:\n            self.scaler.step(optimizer)\n            self.scaler.update()\n        else:\n            optimizer.step()\n\n    def prepare(self, model, optimizer):\n        model = model.to(self.device)\n        if self.num_processes > 1:\n            ddp_kwargs = {\n                \"broadcast_buffers\": False,\n                \"gradient_as_bucket_view\": True,\n            }\n            if self.device.type == \"cuda\":\n                ddp_kwargs[\"device_ids\"] = [self.device.index]\n                ddp_kwargs[\"output_device\"] = self.device.index\n            if self.enable_static_graph:\n                try:\n                    model = DDP(model, static_graph=True, **ddp_kwargs)\n                except TypeError:\n                    model = DDP(model, **ddp_kwargs)\n            else:\n                model = DDP(model, **ddp_kwargs)\n        return model, optimizer\n\n    def unwrap_model(self, model):\n        return model.module if isinstance(model, DDP) else model\n\n    def reduce(self, tensor: torch.Tensor, reduction: str = \"sum\") -> torch.Tensor:\n        if self.num_processes <= 1:\n            return tensor\n        out = tensor.clone()\n        reduction_name = str(reduction).lower()\n        if reduction_name == \"sum\":\n            op = dist.ReduceOp.SUM\n            dist.all_reduce(out, op=op)\n        elif reduction_name == \"mean\":\n            dist.all_reduce(out, op=dist.ReduceOp.SUM)\n            out = out / float(self.num_processes)\n        elif reduction_name == \"max\":\n            dist.all_reduce(out, op=dist.ReduceOp.MAX)\n        else:\n            raise ValueError(f\"Unsupported reduction: {reduction}\")\n        return out\n\n    def gather(self, tensor: torch.Tensor) -> torch.Tensor:\n        if self.num_processes <= 1:\n            return tensor\n        gather_list = [torch.empty_like(tensor) for _ in range(self.num_processes)]\n        dist.all_gather(gather_list, tensor)\n        return torch.cat(gather_list, dim=0)\n\n    def wait_for_everyone(self) -> None:\n        if self.num_processes > 1:\n            if self.device.type == \"cuda\":\n                dist.barrier(device_ids=[self.device.index])\n            else:\n                dist.barrier()\n\n    def close(self) -> None:\n        if dist.is_initialized():\n            dist.destroy_process_group()\n\n\ndef build_runtime(runtime_backend: str, use_amp: bool, enable_static_graph: bool = True):\n    backend = str(runtime_backend).strip().lower()\n    if backend == \"native\":\n        return NativeDDPRuntime(use_amp=use_amp, enable_static_graph=enable_static_graph)\n    if backend == \"accelerate\":\n        return AccelerateRuntime(use_amp=use_amp)\n    raise ValueError(f\"Unsupported runtime_backend: {runtime_backend}\")\n\n\n@dataclass\nclass DayStore:\n    name: str\n    n_rows: int\n    n_factors: int\n    x: np.memmap\n    y: np.memmap\n    w: np.memmap\n    sym_start: np.memmap\n    sym_end: np.memmap\n    need_clean: bool\n\n\n_SPLIT_X_ARRAY_CACHE: Dict[str, np.ndarray] = {}\n\n\ndef _resolve_source_x_path(day_dir: Path, meta: dict) -> Path | None:\n    explicit = str(meta.get(\"source_x_path\") or \"\").strip()\n    if explicit:\n        return Path(explicit)\n    source_root = str(meta.get(\"source_root\") or \"\").strip()\n    source_split = str(meta.get(\"source_split\") or \"\").strip()\n    if source_root and source_split:\n        return Path(source_root) / source_split / \"x.npy\"\n    return None\n\n\ndef _load_cached_x_npy(path: Path) -> np.ndarray:\n    key = str(path.resolve())\n    arr = _SPLIT_X_ARRAY_CACHE.get(key)\n    if arr is None:\n        arr = np.load(path, mmap_mode=\"r\", allow_pickle=False)\n        _SPLIT_X_ARRAY_CACHE[key] = arr\n    return arr\n\n\ndef list_ready_days(row_root: Path) -> List[Path]:\n    days = []\n    for p in sorted(row_root.iterdir()):\n        if p.is_dir() and (p / \"_SUCCESS\").exists():\n            days.append(p)\n    return days\n\n\ndef split_days(days: List[Path], ratio: float) -> Tuple[List[Path], List[Path]]:\n    n = len(days)\n    if n < 2:\n        return days, days\n    cut = max(1, min(n - 1, int(n * ratio)))\n    return days[:cut], days[cut:]\n\n\ndef _day_str_from_dir(day_dir: Path) -> str:\n    name = day_dir.name\n    if len(name) >= 8:\n        day = name[:8]\n        if day.isdigit():\n            return day\n    return \"\"\n\n\ndef filter_days_by_date(days: List[Path], start_date: str, end_date: str) -> List[Path]:\n    s = (start_date or \"\").strip()\n    e = (end_date or \"\").strip()\n    if not s and not e:\n        return list(days)\n    if s and (len(s) != 8 or not s.isdigit()):\n        raise ValueError(f\"Invalid start_date: {s}\")\n    if e and (len(e) != 8 or not e.isdigit()):\n        raise ValueError(f\"Invalid end_date: {e}\")\n    lo = s if s else \"00000000\"\n    hi = e if e else \"99999999\"\n    if lo > hi:\n        raise ValueError(f\"Invalid date window: {lo} > {hi}\")\n    out: List[Path] = []\n    for d in days:\n        day = _day_str_from_dir(d)\n        if day and lo <= day <= hi:\n            out.append(d)\n    return out\n\n\ndef limit_days(days: List[Path], day_limit: int) -> List[Path]:\n    if day_limit <= 0 or day_limit >= len(days):\n        return days\n    return days[:day_limit]\n\n\ndef load_day_meta(day_dir: Path) -> dict:\n    return json.loads((day_dir / \"meta.json\").read_text(encoding=\"utf-8\"))\n\n\ndef load_day_store(day_dir: Path, prefer_fp16: bool) -> DayStore:\n    meta = load_day_meta(day_dir)\n    n_rows = int(meta[\"n_rows\"])\n    n_factors = int(meta[\"n_factors\"])\n    source_x_path = _resolve_source_x_path(day_dir, meta)\n    if source_x_path is not None:\n        row_start = int(meta.get(\"source_row_start\", 0))\n        row_stop = int(meta.get(\"source_row_stop\", row_start + n_rows))\n        x_all = _load_cached_x_npy(source_x_path)\n        x = x_all[row_start:row_stop]\n        need_clean = True\n    else:\n        fp16_path = day_dir / \"x_top200_f16_filled.memmap\"\n        if prefer_fp16 and fp16_path.exists():\n            x = np.memmap(fp16_path, mode=\"r\", dtype=np.float16, shape=(n_rows, n_factors))\n            need_clean = False\n        else:\n            x = np.memmap(day_dir / \"x_top200_float32.memmap\", mode=\"r\", dtype=np.float32, shape=(n_rows, n_factors))\n            need_clean = True\n    y = np.memmap(day_dir / \"y_sum_float32.memmap\", mode=\"r\", dtype=np.float32, shape=(n_rows,))\n    w = np.memmap(day_dir / \"w_float32.memmap\", mode=\"r\", dtype=np.float32, shape=(n_rows,))\n    sym_n = int(meta[\"symbol_count_meta\"])\n    sym_start = np.memmap(day_dir / \"symbol_start_idx_int64.memmap\", mode=\"r\", dtype=np.int64, shape=(sym_n,))\n    sym_end = np.memmap(day_dir / \"symbol_end_idx_int64.memmap\", mode=\"r\", dtype=np.int64, shape=(sym_n,))\n    return DayStore(day_dir.name, n_rows, n_factors, x, y, w, sym_start, sym_end, need_clean)\n\n\ndef load_day_symbol_bounds(day_dir: Path) -> Tuple[np.memmap, np.memmap]:\n    meta = load_day_meta(day_dir)\n    sym_n = int(meta[\"symbol_count_meta\"])\n    sym_start = np.memmap(day_dir / \"symbol_start_idx_int64.memmap\", mode=\"r\", dtype=np.int64, shape=(sym_n,))\n    sym_end = np.memmap(day_dir / \"symbol_end_idx_int64.memmap\", mode=\"r\", dtype=np.int64, shape=(sym_n,))\n    return sym_start, sym_end\n\n\ndef build_valid_end_indices_from_bounds(\n    sym_start: np.ndarray,\n    sym_end: np.ndarray,\n    seq_len: int,\n    sample_stride: int,\n) -> np.ndarray:\n    stride = int(sample_stride)\n    starts = np.asarray(sym_start, dtype=np.int64) + int(seq_len) - 1\n    # symbol_end_idx is exclusive in this memmap layout.\n    ends = np.asarray(sym_end, dtype=np.int64) - 1\n    valid_mask = starts <= ends\n    if not np.any(valid_mask):\n        return np.empty((0,), dtype=np.int64)\n    starts = starts[valid_mask]\n    ends = ends[valid_mask]\n    lengths = ((ends - starts) // stride + 1).astype(np.int64, copy=False)\n    total = int(lengths.sum())\n    if total <= 0:\n        return np.empty((0,), dtype=np.int64)\n    repeated_starts = np.repeat(starts, lengths)\n    group_offsets = np.repeat(np.cumsum(lengths, dtype=np.int64) - lengths, lengths)\n    intra_offsets = np.arange(total, dtype=np.int64) - group_offsets\n    return repeated_starts + intra_offsets * stride\n\n\ndef build_valid_end_indices(store: DayStore, seq_len: int, sample_stride: int) -> np.ndarray:\n    return build_valid_end_indices_from_bounds(\n        store.sym_start,\n        store.sym_end,\n        seq_len=seq_len,\n        sample_stride=sample_stride,\n    )\n\n\ndef build_end_index_cache_path(cache_dir: Path, day_dir: Path, seq_len: int, sample_stride: int) -> Path:\n    key_src = \"\\n\".join(\n        [\n            str(day_dir.parent),\n            str(day_dir.name),\n            str(int(seq_len)),\n            str(int(sample_stride)),\n        ]\n    )\n    key = hashlib.sha1(key_src.encode(\"utf-8\")).hexdigest()[:16]\n    return cache_dir / f\"{day_dir.name}_seq{int(seq_len)}_stride{int(sample_stride)}_{key}.npy\"\n\n\ndef load_end_index_cache(path: Path) -> np.ndarray:\n    arr = np.load(path, allow_pickle=False)\n    return np.asarray(arr, dtype=np.int64)\n\n\ndef wait_for_end_index_cache(path: Path, timeout_sec: float = 600.0) -> np.ndarray:\n    deadline = time.time() + float(timeout_sec)\n    last_err = None\n    while time.time() < deadline:\n        if path.exists() and path.stat().st_size > 0:\n            try:\n                return load_end_index_cache(path)\n            except Exception as exc:  # pragma: no cover - transient partial-write case\n                last_err = exc\n        time.sleep(0.2)\n    if last_err is not None:\n        raise TimeoutError(f\"Timed out waiting for end-index cache {path}: {last_err}\") from last_err\n    raise TimeoutError(f\"Timed out waiting for end-index cache {path}\")\n\n\ndef load_or_build_valid_end_indices(\n    day_dir: Path,\n    seq_len: int,\n    sample_stride: int,\n    cache_dir: Path | None,\n    cache_writer: bool,\n) -> np.ndarray:\n    if cache_dir is None:\n        sym_start, sym_end = load_day_symbol_bounds(day_dir)\n        return build_valid_end_indices_from_bounds(sym_start, sym_end, seq_len=seq_len, sample_stride=sample_stride)\n    cache_path = build_end_index_cache_path(cache_dir, day_dir, seq_len=seq_len, sample_stride=sample_stride)\n    if cache_path.exists() and cache_path.stat().st_size > 0:\n        return load_end_index_cache(cache_path)\n    if not cache_writer:\n        return wait_for_end_index_cache(cache_path)\n    sym_start, sym_end = load_day_symbol_bounds(day_dir)\n    ends = build_valid_end_indices_from_bounds(sym_start, sym_end, seq_len=seq_len, sample_stride=sample_stride)\n    tmp_path = cache_path.with_suffix(f\".tmp.{int(time.time() * 1000)}.{os.getpid()}.npy\")\n    np.save(tmp_path, ends)\n    tmp_path.replace(cache_path)\n    return ends\n\n\ndef truncate_end_indices(end_indices: List[np.ndarray], max_samples: int) -> List[np.ndarray]:\n    if max_samples <= 0:\n        return end_indices\n    remaining = int(max_samples)\n    trimmed: List[np.ndarray] = []\n    for ends in end_indices:\n        if remaining <= 0:\n            trimmed.append(np.empty((0,), dtype=np.int64))\n            continue\n        take = min(int(ends.shape[0]), remaining)\n        trimmed.append(ends[:take])\n        remaining -= take\n    return trimmed\n\n\ndef split_contiguous_even(total: int, rank: int, world_size: int) -> Tuple[int, int]:\n    if total <= 0:\n        return 0, 0\n    start = (total * rank) // max(1, world_size)\n    end = (total * (rank + 1)) // max(1, world_size)\n    return int(start), int(end)\n\n\n@dataclass(frozen=True)\nclass BatchSlice:\n    day_i: int\n    start: int\n    stop: int\n\n\n@dataclass(frozen=True)\nclass BatchPlanItem:\n    parts: Tuple[BatchSlice, ...]\n    pad_size: int = 0\n\n\n@dataclass(frozen=True)\nclass PackedSeqBatchPart:\n    span_x: torch.Tensor\n    local_end: torch.Tensor\n    y: torch.Tensor\n    w: torch.Tensor\n\n\n@dataclass(frozen=True)\nclass PackedSeqBatch:\n    parts: Tuple[PackedSeqBatchPart, ...]\n\n\n@dataclass(frozen=True)\nclass MaterializedSeqBatch:\n    xb: torch.Tensor\n    y: torch.Tensor\n    w: torch.Tensor\n\n\nclass SeqMemmapBatchLoader:\n    def __init__(\n        self,\n        day_dirs: List[Path],\n        seq_len: int,\n        sample_stride: int,\n        batch_size: int,\n        shuffle: bool,\n        max_samples: int = -1,\n        prefer_fp16: bool = True,\n        rank: int = 0,\n        world_size: int = 1,\n        seed: int = 20260312,\n        pin_memory: bool = True,\n        loader_threads: int = 1,\n        prefetch_batches: int = 1,\n        pad_last_batch: bool = False,\n        batch_overlap: int = 0,\n        index_cache_dir: Path | None = None,\n    ):\n        self.seq_len = seq_len\n        self.batch_size = int(batch_size)\n        self.batch_overlap = max(0, int(batch_overlap))\n        if self.batch_overlap >= self.batch_size:\n            raise ValueError(\n                f\"batch_overlap must be smaller than batch_size, got overlap={self.batch_overlap} batch_size={self.batch_size}\"\n            )\n        self.batch_stride = self.batch_size - self.batch_overlap\n        self.shuffle = bool(shuffle)\n        self.rank = int(rank)\n        self.world_size = int(world_size)\n        self.seed = int(seed)\n        self.day_dirs = list(day_dirs)\n        self.prefer_fp16 = bool(prefer_fp16)\n        self.pin_memory = bool(pin_memory) and torch.cuda.is_available()\n        self.loader_threads = max(1, int(loader_threads))\n        self.prefetch_batches = max(1, int(prefetch_batches))\n        self.pad_last_batch = bool(pad_last_batch)\n        self.epoch = 0\n        self.index_cache_dir = index_cache_dir\n        self.stores: List[DayStore | None] = [None for _ in self.day_dirs]\n        end_indices = [\n            load_or_build_valid_end_indices(\n                d,\n                seq_len=seq_len,\n                sample_stride=sample_stride,\n                cache_dir=self.index_cache_dir,\n                cache_writer=(self.rank == 0),\n            )\n            for d in self.day_dirs\n        ]\n        self.end_indices: List[np.ndarray] = truncate_end_indices(end_indices, max_samples=max_samples)\n        self.total_samples = int(sum(int(ends.shape[0]) for ends in self.end_indices))\n        self.seq_offsets = np.arange(-(self.seq_len - 1), 1, dtype=np.int64)\n        self.rank_sample_counts = self._compute_rank_sample_counts()\n        self.rank_total_samples = [int(sum(day_counts)) for day_counts in self.rank_sample_counts]\n        self.rank_unique_batches = [\n            self._compute_batch_count(total) for total in self.rank_total_samples\n        ]\n        self.total_batches = max(self.rank_unique_batches, default=0)\n\n    def __len__(self) -> int:\n        return self.total_batches\n\n    def set_epoch(self, epoch: int) -> None:\n        self.epoch = int(epoch)\n\n    def _get_store(self, day_i: int) -> DayStore:\n        store = self.stores[day_i]\n        if store is None:\n            store = load_day_store(self.day_dirs[day_i], prefer_fp16=self.prefer_fp16)\n            self.stores[day_i] = store\n        return store\n\n    def _compute_rank_sample_counts(self) -> List[List[int]]:\n        counts: List[List[int]] = [[] for _ in range(self.world_size)]\n        for ends in self.end_indices:\n            n = int(ends.shape[0])\n            for rank in range(self.world_size):\n                start, stop = split_contiguous_even(n, rank, self.world_size)\n                counts[rank].append(max(0, stop - start))\n        return counts\n\n    def _compute_batch_count(self, total: int) -> int:\n        total = int(total)\n        if total <= 0:\n            return 0\n        if total <= self.batch_size:\n            return 1\n        return 1 + int(math.ceil((total - self.batch_size) / self.batch_stride))\n\n    def _tail_parts(self, parts: List[BatchSlice], keep: int) -> List[BatchSlice]:\n        keep = max(0, int(keep))\n        if keep <= 0:\n            return []\n        out: List[BatchSlice] = []\n        remaining = keep\n        for part in reversed(parts):\n            part_len = int(part.stop - part.start)\n            if part_len <= 0:\n                continue\n            take = min(remaining, part_len)\n            out.append(BatchSlice(day_i=part.day_i, start=part.stop - take, stop=part.stop))\n            remaining -= take\n            if remaining == 0:\n                break\n        if remaining != 0:\n            raise RuntimeError(f\"Failed to preserve batch overlap keep={keep}, remaining={remaining}\")\n        out.reverse()\n        return out\n\n    def _build_rank_plan(self) -> List[BatchPlanItem]:\n        # Rotate the contiguous shard every epoch so each rank sees different regions over time.\n        shard_rank = (self.rank + self.epoch) % max(1, self.world_size) if self.shuffle else self.rank\n        items: List[BatchPlanItem] = []\n        current_parts: List[BatchSlice] = []\n        filled = 0\n        for day_i, ends in enumerate(self.end_indices):\n            total = int(ends.shape[0])\n            start, stop = split_contiguous_even(total, shard_rank, self.world_size)\n            if stop <= start:\n                continue\n            cursor = int(start)\n            stop = int(stop)\n            while cursor < stop:\n                need = self.batch_size - filled\n                take = min(need, stop - cursor)\n                current_parts.append(BatchSlice(day_i=day_i, start=cursor, stop=cursor + take))\n                cursor += take\n                filled += take\n                if filled == self.batch_size:\n                    emitted_parts = tuple(current_parts)\n                    items.append(BatchPlanItem(parts=emitted_parts, pad_size=0))\n                    current_parts = self._tail_parts(list(emitted_parts), self.batch_overlap)\n                    filled = self.batch_overlap\n        if current_parts:\n            pad_size = self.batch_size - filled if self.pad_last_batch else 0\n            items.append(BatchPlanItem(parts=tuple(current_parts), pad_size=pad_size))\n        if not items:\n            return []\n        if self.shuffle and len(items) > 1:\n            rng = np.random.default_rng(self.seed + self.epoch)\n            if self.batch_overlap > 0:\n                shift = int(rng.integers(len(items)))\n                if shift > 0:\n                    items = items[shift:] + items[:shift]\n            else:\n                order = rng.permutation(len(items))\n                items = [items[int(i)] for i in order.tolist()]\n        if len(items) < self.total_batches:\n            base = list(items)\n            pad_idx = 0\n            while len(items) < self.total_batches:\n                items.append(base[pad_idx % len(base)])\n                pad_idx += 1\n        return items\n\n    def _load_batch_part(self, part: BatchSlice) -> PackedSeqBatchPart:\n        day_i, start, stop = part.day_i, part.start, part.stop\n        store = self._get_store(day_i)\n        batch_end = self.end_indices[day_i][start:stop]\n        if batch_end.size == 0:\n            raise RuntimeError(f\"Empty batch slice for day_i={day_i}, start={start}, stop={stop}\")\n        span_start = int(batch_end[0]) - self.seq_len + 1\n        span_end = int(batch_end[-1])\n        span_x_dtype = np.float32 if store.need_clean else store.x.dtype\n        span_x = np.array(store.x[span_start : span_end + 1], dtype=span_x_dtype, copy=True)\n        if store.need_clean:\n            np.nan_to_num(span_x, copy=False, nan=0.0, posinf=0.0, neginf=0.0)\n        local_end = np.ascontiguousarray(batch_end - span_start, dtype=np.int64)\n        y = np.array(store.y[batch_end], dtype=np.float32, copy=True)\n        w = np.array(store.w[batch_end], dtype=np.float32, copy=True)\n        span_xb = torch.from_numpy(np.ascontiguousarray(span_x))\n        local_endb = torch.from_numpy(local_end)\n        yb = torch.from_numpy(np.ascontiguousarray(y))\n        wb = torch.from_numpy(np.ascontiguousarray(w))\n        if self.pin_memory:\n            span_xb = span_xb.pin_memory()\n            local_endb = local_endb.pin_memory()\n            yb = yb.pin_memory()\n            wb = wb.pin_memory()\n        return PackedSeqBatchPart(span_x=span_xb, local_end=local_endb, y=yb, w=wb)\n\n    def _pad_batch_part(self, part: PackedSeqBatchPart, pad_size: int) -> PackedSeqBatchPart:\n        if pad_size <= 0:\n            return part\n        local_end = torch.cat([part.local_end, part.local_end[-1:].repeat(int(pad_size))], dim=0)\n        y = torch.cat([part.y, part.y[-1:].repeat(int(pad_size))], dim=0)\n        w = torch.cat([part.w, part.w[-1:].repeat(int(pad_size))], dim=0)\n        if self.pin_memory:\n            local_end = local_end.pin_memory()\n            y = y.pin_memory()\n            w = w.pin_memory()\n        return PackedSeqBatchPart(span_x=part.span_x, local_end=local_end, y=y, w=w)\n\n    def _load_batch_from_plan_item(self, item: BatchPlanItem) -> PackedSeqBatch:\n        parts = [self._load_batch_part(part) for part in item.parts]\n        if not parts:\n            raise RuntimeError(\"Empty batch plan item\")\n        if item.pad_size > 0:\n            parts[-1] = self._pad_batch_part(parts[-1], item.pad_size)\n        return PackedSeqBatch(parts=tuple(parts))\n\n    def __iter__(self):\n        plan = self._build_rank_plan()\n        if not plan:\n            return\n        if self.loader_threads <= 1 and self.prefetch_batches <= 1:\n            for item in plan:\n                yield self._load_batch_from_plan_item(item)\n            return\n\n        submit_ahead = max(self.prefetch_batches, self.loader_threads)\n        futures: List[Future] = []\n        next_idx = 0\n        with ThreadPoolExecutor(max_workers=self.loader_threads) as executor:\n            while next_idx < len(plan) and len(futures) < submit_ahead:\n                futures.append(executor.submit(self._load_batch_from_plan_item, plan[next_idx]))\n                next_idx += 1\n            while futures:\n                fut = futures.pop(0)\n                yield fut.result()\n                if next_idx < len(plan):\n                    futures.append(executor.submit(self._load_batch_from_plan_item, plan[next_idx]))\n                    next_idx += 1\n\n\ndef move_batch_to_device(\n    batch: PackedSeqBatch,\n    device: torch.device,\n    copy_non_blocking: bool = False,\n) -> PackedSeqBatch:\n    moved_parts = []\n    for part in batch.parts:\n        moved_parts.append(\n            PackedSeqBatchPart(\n                span_x=part.span_x.to(device, non_blocking=copy_non_blocking),\n                local_end=part.local_end.to(device, non_blocking=copy_non_blocking),\n                y=part.y.to(device, non_blocking=copy_non_blocking),\n                w=part.w.to(device, non_blocking=copy_non_blocking),\n            )\n        )\n    return PackedSeqBatch(parts=tuple(moved_parts))\n\n\nclass DeviceTransferPrefetchLoader:\n    def __init__(\n        self,\n        loader,\n        device: torch.device,\n        prefetch_batches: int = 2,\n        copy_non_blocking: bool = False,\n    ):\n        self.loader = loader\n        self.device = device\n        self.prefetch_batches = max(1, int(prefetch_batches))\n        self.copy_non_blocking = bool(copy_non_blocking)\n\n    def __len__(self) -> int:\n        return len(self.loader)\n\n    def __getattr__(self, name: str):\n        return getattr(self.loader, name)\n\n    def set_epoch(self, epoch: int) -> None:\n        if hasattr(self.loader, \"set_epoch\"):\n            self.loader.set_epoch(epoch)\n\n    def __iter__(self):\n        if self.device.type != \"cuda\":\n            for batch in self.loader:\n                yield batch\n            return\n        result_q: queue.Queue = queue.Queue(maxsize=self.prefetch_batches)\n        sentinel = object()\n        device = self.device\n        stop_event = threading.Event()\n\n        def put_result(item, event) -> bool:\n            while not stop_event.is_set():\n                try:\n                    result_q.put((item, event), timeout=0.1)\n                    return True\n                except queue.Full:\n                    continue\n            return False\n\n        def worker():\n            try:\n                torch.cuda.set_device(device)\n                stream = torch.cuda.Stream(device=device)\n                for batch in self.loader:\n                    if stop_event.is_set():\n                        break\n                    with torch.cuda.stream(stream):\n                        moved = move_batch_to_device(\n                            batch,\n                            device=device,\n                            copy_non_blocking=self.copy_non_blocking,\n                        )\n                        event = torch.cuda.Event()\n                        event.record(stream)\n                    if not put_result(moved, event):\n                        return\n                put_result(sentinel, None)\n            except Exception as exc:  # pragma: no cover - worker thread failure propagation\n                put_result(exc, None)\n\n        thread = threading.Thread(target=worker, daemon=True)\n        thread.start()\n        current_stream = torch.cuda.current_stream(device)\n        try:\n            while True:\n                item, event = result_q.get()\n                if item is sentinel:\n                    break\n                if isinstance(item, Exception):\n                    raise item\n                current_stream.wait_event(event)\n                for part in item.parts:\n                    part.span_x.record_stream(current_stream)\n                    part.local_end.record_stream(current_stream)\n                    part.y.record_stream(current_stream)\n                    part.w.record_stream(current_stream)\n                yield item\n        finally:\n            stop_event.set()\n            thread.join()\n\n\nclass MaterializeDevicePrefetchLoader:\n    def __init__(\n        self,\n        loader,\n        device: torch.device,\n        seq_offsets: torch.Tensor,\n        feat_mean: torch.Tensor,\n        feat_std: torch.Tensor,\n        prefetch_batches: int = 2,\n        copy_non_blocking: bool = False,\n    ):\n        self.loader = loader\n        self.device = device\n        self.seq_offsets = seq_offsets\n        self.feat_mean = feat_mean\n        self.feat_std = feat_std\n        self.prefetch_batches = max(1, int(prefetch_batches))\n        self.copy_non_blocking = bool(copy_non_blocking)\n\n    def __len__(self) -> int:\n        return len(self.loader)\n\n    def __getattr__(self, name: str):\n        return getattr(self.loader, name)\n\n    def set_epoch(self, epoch: int) -> None:\n        if hasattr(self.loader, \"set_epoch\"):\n            self.loader.set_epoch(epoch)\n\n    def __iter__(self):\n        if self.device.type != \"cuda\":\n            for batch in self.loader:\n                xb, yb, wb = materialize_batch(\n                    batch,\n                    seq_offsets=self.seq_offsets,\n                    feat_mean=self.feat_mean,\n                    feat_std=self.feat_std,\n                    device=self.device,\n                    copy_non_blocking=self.copy_non_blocking,\n                )\n                yield MaterializedSeqBatch(xb=xb, y=yb, w=wb)\n            return\n        result_q: queue.Queue = queue.Queue(maxsize=self.prefetch_batches)\n        sentinel = object()\n        device = self.device\n        stop_event = threading.Event()\n\n        def put_result(item, event) -> bool:\n            while not stop_event.is_set():\n                try:\n                    result_q.put((item, event), timeout=0.1)\n                    return True\n                except queue.Full:\n                    continue\n            return False\n\n        def worker():\n            try:\n                torch.cuda.set_device(device)\n                stream = torch.cuda.Stream(device=device)\n                for batch in self.loader:\n                    if stop_event.is_set():\n                        break\n                    with torch.cuda.stream(stream):\n                        xb, yb, wb = materialize_batch(\n                            batch,\n                            seq_offsets=self.seq_offsets,\n                            feat_mean=self.feat_mean,\n                            feat_std=self.feat_std,\n                            device=device,\n                            copy_non_blocking=self.copy_non_blocking,\n                        )\n                        event = torch.cuda.Event()\n                        event.record(stream)\n                    prefetched = MaterializedSeqBatch(xb=xb, y=yb, w=wb)\n                    if not put_result(prefetched, event):\n                        return\n                put_result(sentinel, None)\n            except Exception as exc:  # pragma: no cover - worker thread failure propagation\n                put_result(exc, None)\n\n        thread = threading.Thread(target=worker, daemon=True)\n        thread.start()\n        current_stream = torch.cuda.current_stream(device)\n        try:\n            while True:\n                item, event = result_q.get()\n                if item is sentinel:\n                    break\n                if isinstance(item, Exception):\n                    raise item\n                current_stream.wait_event(event)\n                item.xb.record_stream(current_stream)\n                item.y.record_stream(current_stream)\n                item.w.record_stream(current_stream)\n                yield item\n        finally:\n            stop_event.set()\n            thread.join()\n\n\nclass ParallelCNN1D(nn.Module):\n    def __init__(self, in_channels: int, out_channels: int):\n        super().__init__()\n        branch_channels = out_channels // 3\n        self.conv3 = nn.Conv1d(in_channels, branch_channels, kernel_size=3, padding=1)\n        self.conv5 = nn.Conv1d(in_channels, branch_channels, kernel_size=5, padding=2)\n        self.conv7 = nn.Conv1d(in_channels, out_channels - 2 * branch_channels, kernel_size=7, padding=3)\n        self.act = nn.ReLU()\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x_t = x.transpose(1, 2)\n        c3 = self.conv3(x_t)\n        c5 = self.conv5(x_t)\n        c7 = self.conv7(x_t)\n        out = torch.cat([c3, c5, c7], dim=1)\n        return self.act(out).transpose(1, 2)\n\n\nclass FeatureChannelGate(nn.Module):\n    def __init__(self, input_dim: int, hidden_dim: int = 64, init_bias: float = 2.0):\n        super().__init__()\n        hidden_dim = max(1, int(hidden_dim))\n        self.norm = nn.LayerNorm(input_dim)\n        self.fc1 = nn.Linear(input_dim, hidden_dim)\n        self.act = nn.SiLU()\n        self.fc2 = nn.Linear(hidden_dim, input_dim)\n        nn.init.zeros_(self.fc2.weight)\n        nn.init.constant_(self.fc2.bias, float(init_bias))\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        pooled = x.mean(dim=1)\n        gate = self.fc2(self.act(self.fc1(self.norm(pooled))))\n        gate = torch.sigmoid(gate).unsqueeze(1)\n        return x * gate\n\n\nclass GRURegressor(nn.Module):\n    def __init__(\n        self,\n        input_dim: int,\n        hidden_dim: int,\n        num_layers: int,\n        dropout: float,\n        pooling: str = \"last\",\n        bidirectional: bool = False,\n        use_cnn1d: bool = False,\n        input_gate_hidden_dim: int = 0,\n        input_gate_bias: float = 2.0,\n    ):\n        super().__init__()\n        self.pooling = str(pooling).lower()\n        if self.pooling not in {\"last\", \"attn\"}:\n            raise ValueError(f\"Unsupported pooling: {pooling}\")\n        self.bidirectional = bool(bidirectional)\n        self.use_cnn1d = bool(use_cnn1d)\n        self.input_gate_hidden_dim = max(0, int(input_gate_hidden_dim))\n        self.output_dim = int(hidden_dim) * (2 if self.bidirectional else 1)\n        if self.input_gate_hidden_dim > 0:\n            self.input_gate = FeatureChannelGate(\n                input_dim=input_dim,\n                hidden_dim=self.input_gate_hidden_dim,\n                init_bias=float(input_gate_bias),\n            )\n        else:\n            self.input_gate = nn.Identity()\n\n        if self.use_cnn1d:\n            self.cnn = ParallelCNN1D(input_dim, hidden_dim)\n            rnn_input_dim = hidden_dim\n        else:\n            self.cnn = nn.Identity()\n            rnn_input_dim = input_dim\n\n        self.rnn = nn.GRU(\n            input_size=rnn_input_dim,\n            hidden_size=hidden_dim,\n            num_layers=num_layers,\n            dropout=dropout if num_layers > 1 else 0.0,\n            batch_first=True,\n            bidirectional=self.bidirectional,\n        )\n        if self.pooling == \"attn\":\n            self.attn_norm = nn.LayerNorm(self.output_dim)\n            self.attn_proj = nn.Linear(self.output_dim, 1)\n        head_hidden_dim = max(1, self.output_dim // 2)\n        self.head = nn.Sequential(\n            nn.LayerNorm(self.output_dim),\n            nn.Linear(self.output_dim, head_hidden_dim),\n            nn.ReLU(),\n            nn.Linear(head_hidden_dim, 1),\n        )\n\n    def forward(self, x: torch.Tensor) -> torch.Tensor:\n        x = self.input_gate(x)\n        if self.use_cnn1d:\n            x = self.cnn(x)\n        out, _ = self.rnn(x)\n        if self.pooling == \"attn\":\n            score = self.attn_proj(self.attn_norm(out)).squeeze(-1)\n            weight = torch.softmax(score, dim=1).unsqueeze(-1)\n            pooled = (out * weight).sum(dim=1)\n        else:\n            pooled = out[:, -1, :]\n        return self.head(pooled).squeeze(-1)\n\n\ndef weighted_mse(pred: torch.Tensor, y: torch.Tensor, w: torch.Tensor) -> torch.Tensor:\n    w = torch.clamp(w, min=0.0)\n    w_norm = w / torch.clamp(w.mean(), min=1e-6)\n    return (w_norm * (pred - y) ** 2).mean()\n\n\ndef build_val_epoch_checkpoint_path(out_root: Path, epoch: int) -> Path:\n    return out_root / f\"gru_seq_memmap_val_epoch_{int(epoch):03d}.pt\"\n\n\ndef refresh_top_val_checkpoints(\n    top_val_checkpoints: List[dict],\n    candidate_epoch: int,\n    candidate_val_metrics: dict,\n    keep_topk: int,\n    out_root: Path,\n) -> Tuple[List[dict], bool, List[dict]]:\n    keep_topk = max(1, int(keep_topk))\n    candidate = {\n        \"epoch\": int(candidate_epoch),\n        \"val_ic\": float(candidate_val_metrics[\"ic\"]),\n        \"val_loss\": float(candidate_val_metrics[\"loss\"]),\n        \"path\": str(build_val_epoch_checkpoint_path(out_root, candidate_epoch)),\n    }\n    updated = list(top_val_checkpoints)\n    updated.append(candidate)\n    updated.sort(key=lambda item: (-float(item[\"val_ic\"]), int(item[\"epoch\"])))\n    kept = [dict(item) for item in updated[:keep_topk]]\n    dropped = [dict(item) for item in updated[keep_topk:]]\n    kept_epochs = {int(item[\"epoch\"]) for item in kept}\n    entered_topk = int(candidate_epoch) in kept_epochs\n    for rank_i, item in enumerate(kept, start=1):\n        item[\"rank\"] = int(rank_i)\n    return kept, entered_topk, dropped\n\n\nclass CorrStats:\n    def __init__(self):\n        self.buf: torch.Tensor | None = None\n\n    def update(self, pred: torch.Tensor, y: torch.Tensor, w: torch.Tensor | None = None):\n        p = pred.detach().float().reshape(-1)\n        t = y.detach().float().reshape(-1)\n        mask = torch.isfinite(p) & torch.isfinite(t)\n        if w is not None:\n            ww = w.detach().float().reshape(-1)\n            mask = mask & torch.isfinite(ww) & (ww > 0)\n        zero = torch.zeros_like(p)\n        p = torch.where(mask, p, zero).to(dtype=torch.float64)\n        t = torch.where(mask, t, zero).to(dtype=torch.float64)\n        d = p - t\n        sums = torch.stack(\n            [\n                mask.to(dtype=torch.float64).sum(),\n                p.sum(),\n                t.sum(),\n                (p * p).sum(),\n                (t * t).sum(),\n                (p * t).sum(),\n                torch.abs(d).sum(),\n                (d * d).sum(),\n            ]\n        )\n        if self.buf is None:\n            self.buf = sums\n        else:\n            self.buf = self.buf + sums\n\n    def to_tensor(self, device: torch.device) -> torch.Tensor:\n        if self.buf is None:\n            return torch.zeros((8,), dtype=torch.float64, device=device)\n        return self.buf.to(device=device, dtype=torch.float64)\n\n\ndef metrics_from_tensor(t: torch.Tensor) -> dict:\n    n, sum_p, sum_y, sum_pp, sum_yy, sum_py, sum_abs, sum_sq = [float(x) for x in t.tolist()]\n    if (not np.isfinite(n)) or n <= 0:\n        return {\"rmse\": float(\"nan\"), \"mae\": float(\"nan\"), \"ic\": float(\"nan\"), \"n\": 0}\n    mse = sum_sq / n\n    var_p = max(sum_pp / n - (sum_p / n) ** 2, 1e-12)\n    var_y = max(sum_yy / n - (sum_y / n) ** 2, 1e-12)\n    cov = sum_py / n - (sum_p / n) * (sum_y / n)\n    ic = cov / math.sqrt(var_p * var_y)\n    return {\"rmse\": float(math.sqrt(mse)), \"mae\": float(sum_abs / n), \"ic\": float(ic), \"n\": int(n)}\n\n\ndef compute_feature_stats(\n    train_days: List[Path],\n    prefer_fp16: bool,\n    sample_stride: int = 50,\n    chunk_sample_rows: int = 250_000,\n) -> Tuple[np.ndarray, np.ndarray]:\n    s = None\n    s2 = None\n    n = 0\n    for day in train_days:\n        store = load_day_store(day, prefer_fp16=prefer_fp16)\n        chunk_span = max(int(sample_stride), int(sample_stride) * max(1, int(chunk_sample_rows)))\n        for start in range(0, store.n_rows, chunk_span):\n            stop = min(store.n_rows, start + chunk_span)\n            # Stream smaller memmap slices to reduce pressure and avoid giant one-shot reads.\n            arr = np.array(store.x[start:stop:sample_stride], dtype=np.float32, copy=True)\n            if arr.size == 0:\n                continue\n            arr = np.nan_to_num(arr, nan=0.0, posinf=0.0, neginf=0.0)\n            if s is None:\n                s = arr.sum(axis=0, dtype=np.float64)\n                s2 = (arr * arr).sum(axis=0, dtype=np.float64)\n            else:\n                s += arr.sum(axis=0, dtype=np.float64)\n                s2 += (arr * arr).sum(axis=0, dtype=np.float64)\n            n += arr.shape[0]\n    mean = (s / max(n, 1)).astype(np.float32)\n    var = (s2 / max(n, 1) - mean.astype(np.float64) ** 2).astype(np.float32)\n    std = np.sqrt(np.clip(var, 1e-8, None)).astype(np.float32)\n    return mean, std\n\n\ndef build_feature_stats_cache_path(\n    output_root: Path,\n    row_root: Path,\n    train_days: List[Path],\n    prefer_fp16: bool,\n    sample_stride: int,\n) -> Path:\n    cache_dir = ensure_dir(output_root / \"training_seq\" / \"_feature_stats_cache\")\n    key_src = \"\\n\".join(\n        [\n            str(row_root),\n            str(bool(prefer_fp16)),\n            str(int(sample_stride)),\n            *[d.name for d in train_days],\n        ]\n    )\n    key = hashlib.sha1(key_src.encode(\"utf-8\")).hexdigest()[:16]\n    return cache_dir / f\"feature_stats_{key}.npz\"\n\n\ndef load_feature_stats_cache(path: Path) -> Tuple[np.ndarray, np.ndarray]:\n    with np.load(path) as arr:\n        mean = arr[\"mean\"].astype(np.float32, copy=False)\n        std = arr[\"std\"].astype(np.float32, copy=False)\n    return mean, std\n\n\ndef wait_for_feature_stats_cache(path: Path, timeout_seconds: int = 7200) -> Tuple[np.ndarray, np.ndarray]:\n    deadline = time.time() + max(1, int(timeout_seconds))\n    last_err: Exception | None = None\n    while time.time() < deadline:\n        if path.exists() and path.stat().st_size > 0:\n            try:\n                return load_feature_stats_cache(path)\n            except Exception as exc:  # pragma: no cover - transient partial-write case\n                last_err = exc\n        time.sleep(2.0)\n    if last_err is not None:\n        raise TimeoutError(f\"Timed out waiting for feature stats cache {path}: {last_err}\") from last_err\n    raise TimeoutError(f\"Timed out waiting for feature stats cache {path}\")\n\n\ndef evaluate(\n    model,\n    loader,\n    feat_mean,\n    feat_std,\n    seq_offsets,\n    accelerator,\n    split_name: str = \"eval\",\n    progress_parts: int = 4,\n    copy_non_blocking: bool = False,\n    collect_timing: bool = False,\n) -> dict:\n    model.eval()\n    stats = CorrStats()\n    loss_sum = torch.zeros((), dtype=torch.float64, device=accelerator.device)\n    loss_count = torch.zeros((), dtype=torch.float64, device=accelerator.device)\n    timing_labels = [\n        \"batch_wait_sec\",\n        \"materialize_sec\",\n        \"forward_sec\",\n        \"stats_update_sec\",\n        \"step_total_sec\",\n    ]\n    timing_sums = np.zeros((len(timing_labels),), dtype=np.float64)\n    local_batches = 0\n    total_batches = 0\n    try:\n        total_batches = len(loader)\n    except TypeError:\n        total_batches = 0\n    progress_step = max(1, total_batches // max(1, int(progress_parts))) if total_batches > 0 else 0\n    with torch.no_grad():\n        loader_iter = iter(loader)\n        batch_i = 0\n        while True:\n            t_step0 = time.perf_counter()\n            t0 = time.perf_counter()\n            try:\n                batch = next(loader_iter)\n            except StopIteration:\n                break\n            t1 = time.perf_counter()\n            xb, yb, wb = resolve_model_batch(\n                batch,\n                seq_offsets=seq_offsets,\n                feat_mean=feat_mean,\n                feat_std=feat_std,\n                device=accelerator.device,\n                copy_non_blocking=copy_non_blocking,\n            )\n            if collect_timing:\n                sync_if_needed(accelerator.device)\n            t2 = time.perf_counter()\n            with accelerator.autocast():\n                pred = model(xb)\n                loss = weighted_mse(pred, yb, wb)\n            if collect_timing:\n                sync_if_needed(accelerator.device)\n            t3 = time.perf_counter()\n            stats.update(pred, yb, wb)\n            if collect_timing:\n                sync_if_needed(accelerator.device)\n            t4 = time.perf_counter()\n            loss_sum = loss_sum + loss.detach().to(dtype=torch.float64)\n            loss_count = loss_count + 1.0\n            local_batches += 1\n            batch_i += 1\n            if collect_timing:\n                timing_sums += np.asarray(\n                    [\n                        t1 - t0,\n                        t2 - t1,\n                        t3 - t2,\n                        t4 - t3,\n                        t4 - t_step0,\n                    ],\n                    dtype=np.float64,\n                )\n            if (\n                accelerator.is_main_process\n                and total_batches > 0\n                and (batch_i % progress_step == 0 or batch_i == total_batches)\n            ):\n                print(f\"[{split_name}] progress {batch_i}/{total_batches}\", flush=True)\n    loss_pair = accelerator.reduce(torch.stack([loss_sum, loss_count]), reduction=\"sum\")\n    loss_den = float(loss_pair[1].item())\n    loss_val = float((loss_pair[0] / loss_den).item()) if loss_den > 0 else float(\"nan\")\n    stats_t = accelerator.reduce(stats.to_tensor(accelerator.device), reduction=\"sum\")\n    m = metrics_from_tensor(stats_t)\n    m[\"loss\"] = float(loss_val)\n    if collect_timing:\n        m[\"_timing\"] = summarize_timing_across_ranks(\n            accelerator=accelerator,\n            labels=timing_labels,\n            count=local_batches,\n            sums=timing_sums,\n            count_key=\"measured_batches\",\n            global_batch=int(loader.batch_size) * int(accelerator.num_processes),\n        )\n    return m\n\n\ndef materialize_batch(\n    batch: PackedSeqBatch,\n    seq_offsets: torch.Tensor,\n    feat_mean: torch.Tensor,\n    feat_std: torch.Tensor,\n    device: torch.device,\n    copy_non_blocking: bool = False,\n) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:\n    xb_parts = []\n    y_parts = []\n    w_parts = []\n    for part in batch.parts:\n        span_x = part.span_x.to(device, non_blocking=copy_non_blocking)\n        local_end = part.local_end.to(device, non_blocking=copy_non_blocking)\n        yb = part.y.to(device, non_blocking=copy_non_blocking)\n        wb = part.w.to(device, non_blocking=copy_non_blocking)\n        seq_idx = local_end[:, None] + seq_offsets[None, :]\n        xb = span_x[seq_idx].float()\n        xb = (xb - feat_mean) / feat_std\n        xb_parts.append(xb)\n        y_parts.append(yb)\n        w_parts.append(wb)\n    if len(xb_parts) == 1:\n        return xb_parts[0], y_parts[0], w_parts[0]\n    return torch.cat(xb_parts, dim=0), torch.cat(y_parts, dim=0), torch.cat(w_parts, dim=0)\n\n\ndef resolve_model_batch(\n    batch,\n    seq_offsets: torch.Tensor,\n    feat_mean: torch.Tensor,\n    feat_std: torch.Tensor,\n    device: torch.device,\n    copy_non_blocking: bool = False,\n) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:\n    if isinstance(batch, MaterializedSeqBatch):\n        return batch.xb, batch.y, batch.w\n    return materialize_batch(\n        batch,\n        seq_offsets=seq_offsets,\n        feat_mean=feat_mean,\n        feat_std=feat_std,\n        device=device,\n        copy_non_blocking=copy_non_blocking,\n    )\n\n\ndef sync_if_needed(device: torch.device) -> None:\n    if device.type == \"cuda\":\n        torch.cuda.synchronize(device)\n\n\ndef warmup_collective(accelerator) -> float:\n    if int(getattr(accelerator, \"num_processes\", 1)) <= 1:\n        return 0.0\n    t0 = time.perf_counter()\n    dummy = torch.zeros((1,), dtype=torch.float32, device=accelerator.device)\n    _ = accelerator.reduce(dummy, reduction=\"sum\")\n    sync_if_needed(accelerator.device)\n    return float(time.perf_counter() - t0)\n\n\ndef summarize_timing_across_ranks(\n    accelerator: Accelerator,\n    labels: List[str],\n    count: int,\n    sums: np.ndarray,\n    count_key: str,\n    global_batch: int | None = None,\n) -> dict:\n    local = torch.tensor([float(count), *sums.tolist()], dtype=torch.float64, device=accelerator.device)[None, :]\n    gathered = accelerator.gather(local)\n    empty = {\n        count_key: int(count),\n        \"per_rank\": [],\n        \"global_step_sec_estimate\": float(\"nan\"),\n        \"global_samples_per_sec_estimate\": float(\"nan\"),\n    }\n    if not accelerator.is_main_process:\n        return empty\n    gathered_np = gathered.detach().cpu().numpy()\n    per_rank = []\n    max_step_sec = 0.0\n    step_label = labels[-1] if labels else None\n    for rank_i, row in enumerate(gathered_np):\n        rank_count = max(1.0, float(row[0]))\n        stats = {\"rank\": int(rank_i), count_key: int(row[0])}\n        for label_i, label in enumerate(labels, start=1):\n            stats[label] = float(row[label_i] / rank_count)\n        if step_label is not None:\n            max_step_sec = max(max_step_sec, float(stats[step_label]))\n        per_rank.append(stats)\n    out = {\n        count_key: int(count),\n        \"per_rank\": per_rank,\n        \"global_step_sec_estimate\": float(max_step_sec) if step_label is not None else float(\"nan\"),\n        \"global_samples_per_sec_estimate\": float(\"nan\"),\n    }\n    if step_label is not None and global_batch is not None:\n        out[\"global_samples_per_sec_estimate\"] = float(global_batch / max(max_step_sec, 1e-9))\n    return out\n\n\ndef profile_train_steps(\n    model,\n    optimizer,\n    loader,\n    accelerator,\n    feat_mean,\n    feat_std,\n    seq_offsets,\n    copy_non_blocking: bool,\n    warmup_steps: int,\n    profile_steps: int,\n    grad_accum_steps: int = 1,\n) -> dict:\n    total_batches = 0\n    try:\n        total_batches = len(loader)\n    except TypeError:\n        total_batches = 0\n    warmup = max(0, int(warmup_steps))\n    active = max(0, int(profile_steps))\n    if total_batches <= 0 or active <= 0:\n        return {\n            \"warmup_steps\": warmup,\n            \"profile_steps\": active,\n            \"measured_steps\": 0,\n            \"per_rank\": [],\n            \"global_step_sec_estimate\": float(\"nan\"),\n            \"global_samples_per_sec_estimate\": float(\"nan\"),\n        }\n    run_steps = min(total_batches, warmup + active)\n    measured_steps = max(0, run_steps - warmup)\n    labels = [\n        \"batch_wait_sec\",\n        \"materialize_sec\",\n        \"forward_sec\",\n        \"backward_sec\",\n        \"optimizer_step_sec\",\n        \"step_total_sec\",\n    ]\n    sums = np.zeros((len(labels),), dtype=np.float64)\n    model.train()\n    train_iter = iter(loader)\n    warmup_collective(accelerator)\n    if accelerator.device.type == \"cuda\":\n        sync_if_needed(accelerator.device)\n        torch.cuda.reset_peak_memory_stats(accelerator.device)\n    accum_steps = max(1, int(grad_accum_steps))\n    accum_micro = 0\n    optimizer.zero_grad(set_to_none=True)\n    for step_i in range(1, run_steps + 1):\n        if accelerator.device.type == \"cuda\" and step_i == warmup + 1:\n            sync_if_needed(accelerator.device)\n            torch.cuda.reset_peak_memory_stats(accelerator.device)\n        t_step0 = time.perf_counter()\n        t0 = time.perf_counter()\n        batch = next(train_iter)\n        t1 = time.perf_counter()\n        accum_micro += 1\n        last_batch = step_i == run_steps\n        should_step = accum_micro >= accum_steps or last_batch\n        group_total = min(accum_steps, accum_micro + max(0, run_steps - step_i))\n        xb, yb, wb = resolve_model_batch(\n            batch,\n            seq_offsets=seq_offsets,\n            feat_mean=feat_mean,\n            feat_std=feat_std,\n            device=accelerator.device,\n            copy_non_blocking=copy_non_blocking,\n        )\n        sync_if_needed(accelerator.device)\n        t2 = time.perf_counter()\n        sync_ctx = model.no_sync() if hasattr(model, \"no_sync\") and not should_step else nullcontext()\n        with sync_ctx:\n            with accelerator.autocast():\n                pred = model(xb)\n                loss_raw = weighted_mse(pred, yb, wb)\n                loss = loss_raw / float(group_total)\n            sync_if_needed(accelerator.device)\n            t3 = time.perf_counter()\n            accelerator.backward(loss)\n            sync_if_needed(accelerator.device)\n            t4 = time.perf_counter()\n        if should_step:\n            accelerator.step_optimizer(optimizer)\n            optimizer.zero_grad(set_to_none=True)\n            accum_micro = 0\n        sync_if_needed(accelerator.device)\n        t5 = time.perf_counter()\n        if step_i > warmup:\n            sums += np.asarray(\n                [\n                    t1 - t0,\n                    t2 - t1,\n                    t3 - t2,\n                    t4 - t3,\n                    t5 - t4,\n                    t5 - t_step0,\n                ],\n                dtype=np.float64,\n            )\n        if accelerator.is_main_process and (step_i == run_steps or step_i % max(1, run_steps // 4) == 0):\n            print(f\"[profile-train] step {step_i}/{run_steps} loss={loss_raw.detach().float().item():.6f}\", flush=True)\n    peak_allocated_bytes = 0.0\n    peak_reserved_bytes = 0.0\n    if accelerator.device.type == \"cuda\":\n        sync_if_needed(accelerator.device)\n        peak_allocated_bytes = float(torch.cuda.max_memory_allocated(accelerator.device))\n        peak_reserved_bytes = float(torch.cuda.max_memory_reserved(accelerator.device))\n    global_batch = int(loader.batch_size) * int(accelerator.num_processes)\n    summary = summarize_timing_across_ranks(\n        accelerator=accelerator,\n        labels=labels,\n        count=measured_steps,\n        sums=sums,\n        count_key=\"measured_steps\",\n        global_batch=global_batch,\n    )\n    summary[\"warmup_steps\"] = warmup\n    summary[\"profile_steps\"] = active\n    if accelerator.device.type == \"cuda\":\n        local_mem = torch.tensor(\n            [peak_allocated_bytes, peak_reserved_bytes],\n            dtype=torch.float64,\n            device=accelerator.device,\n        )[None, :]\n        gathered_mem = accelerator.gather(local_mem)\n        if accelerator.is_main_process:\n            gathered_mem_np = gathered_mem.detach().cpu().numpy()\n            peak_allocated_max = 0.0\n            peak_reserved_max = 0.0\n            for rank_i, row in enumerate(gathered_mem_np):\n                allocated_bytes = float(row[0])\n                reserved_bytes = float(row[1])\n                peak_allocated_max = max(peak_allocated_max, allocated_bytes)\n                peak_reserved_max = max(peak_reserved_max, reserved_bytes)\n                if rank_i < len(summary.get(\"per_rank\", [])):\n                    summary[\"per_rank\"][rank_i][\"peak_allocated_bytes\"] = allocated_bytes\n                    summary[\"per_rank\"][rank_i][\"peak_reserved_bytes\"] = reserved_bytes\n                    summary[\"per_rank\"][rank_i][\"peak_allocated_gib\"] = allocated_bytes / float(1024**3)\n                    summary[\"per_rank\"][rank_i][\"peak_reserved_gib\"] = reserved_bytes / float(1024**3)\n            summary[\"peak_allocated_bytes_max\"] = peak_allocated_max\n            summary[\"peak_reserved_bytes_max\"] = peak_reserved_max\n            summary[\"peak_allocated_gib_max\"] = peak_allocated_max / float(1024**3)\n            summary[\"peak_reserved_gib_max\"] = peak_reserved_max / float(1024**3)\n    return summary\n\n\ndef run(args):\n    run_t0 = time.perf_counter()\n    phase_timings: Dict[str, float] = {}\n    phase_details: Dict[str, dict] = {}\n    data_prepare_t0 = time.perf_counter()\n    set_global_seed(args.seed)\n    cfg = load_config(args.config)\n    train_cfg = dict(cfg[\"train\"])\n    if args.override_epochs > 0:\n        train_cfg[\"epochs\"] = args.override_epochs\n    if args.override_lr > 0:\n        train_cfg[\"learning_rate\"] = args.override_lr\n    if args.override_batch_size > 0:\n        train_cfg[\"batch_size\"] = args.override_batch_size\n\n    row_root = resolve_row_root(cfg, args.cache_name, row_root_override=args.row_root_override)\n    profile_phases = bool(args.profile_phases)\n    profile_warmup_steps = max(0, int(args.profile_warmup_steps))\n    profile_steps = max(0, int(args.profile_steps))\n    stop_after_profile = bool(args.stop_after_profile)\n    profile_only = bool(stop_after_profile and profile_steps > 0)\n    grad_accum_steps = max(1, int(args.grad_accum_steps))\n    split_view_meta = load_split_view_meta(row_root)\n    if split_view_meta is not None:\n        train_days = resolve_named_day_dirs(row_root, split_view_meta.get(\"train_days\"))\n        valid_days = resolve_named_day_dirs(row_root, split_view_meta.get(\"valid_days\"))\n        test_days = resolve_named_day_dirs(row_root, split_view_meta.get(\"test_days\"))\n        train_days = limit_days(train_days, args.train_day_limit)\n        valid_days = limit_days(valid_days, args.valid_day_limit)\n        test_days = limit_days(test_days, args.test_day_limit)\n        train_pool_days = list(train_days) + list(valid_days)\n        days = list(train_days) + list(valid_days) + list(test_days)\n        split_ratio = float(args.valid_split_ratio) if args.valid_split_ratio > 0.0 else float(train_cfg[\"train_split_ratio\"])\n        train_cfg[\"train_split_ratio\"] = split_ratio\n    else:\n        days = list_ready_days(row_root)\n        use_explicit_split = bool(\n            args.train_pool_start_date or args.train_pool_end_date or args.test_start_date or args.test_end_date\n        )\n        if args.max_days > 0 and not use_explicit_split:\n            days = days[: args.max_days]\n        if args.test_start_date or args.test_end_date:\n            test_days = filter_days_by_date(days, args.test_start_date, args.test_end_date)\n        else:\n            test_days = []\n        test_day_names = {d.name for d in test_days}\n        raw_train_pool_days = filter_days_by_date(days, args.train_pool_start_date, args.train_pool_end_date)\n        if args.train_pool_start_date or args.train_pool_end_date:\n            train_pool_days = [d for d in raw_train_pool_days if d.name not in test_day_names]\n        else:\n            train_pool_days = [d for d in days if d.name not in test_day_names]\n        split_ratio = float(args.valid_split_ratio) if args.valid_split_ratio > 0.0 else float(train_cfg[\"train_split_ratio\"])\n        train_cfg[\"train_split_ratio\"] = split_ratio\n        train_days, valid_days = split_days(train_pool_days, split_ratio)\n        train_days = limit_days(train_days, args.train_day_limit)\n        valid_days = limit_days(valid_days, args.valid_day_limit)\n        test_days = limit_days(test_days, args.test_day_limit)\n    if not train_pool_days:\n        raise RuntimeError(\"train_pool_days is empty after applying date filters.\")\n    if len(train_days) == 0 or len(valid_days) == 0:\n        raise RuntimeError(\"Need both train and valid day split.\")\n    print(\n        f\"[data] days_total={len(days)} train_pool_days={len(train_pool_days)} \"\n        f\"train_days={len(train_days)} valid_days={len(valid_days)} test_days={len(test_days)}\",\n        flush=True,\n    )\n    print(f\"[row-root] {row_root}\", flush=True)\n    if args.train_pool_start_date or args.train_pool_end_date:\n        print(\n            f\"[data] train_pool_date_window=[{args.train_pool_start_date or 'min'}, \"\n            f\"{args.train_pool_end_date or 'max'}]\",\n            flush=True,\n        )\n    if args.test_start_date or args.test_end_date:\n        print(\n            f\"[data] test_date_window=[{args.test_start_date or 'min'}, {args.test_end_date or 'max'}]\",\n            flush=True,\n        )\n    data_prepare_local_sec = time.perf_counter() - data_prepare_t0\n    accelerator_init_t0 = time.perf_counter()\n    accelerator = build_runtime(\n        args.runtime_backend,\n        bool(args.use_amp),\n        enable_static_graph=(grad_accum_steps == 1),\n    )\n    phase_timings[\"data_prepare_sec\"] = float(data_prepare_local_sec)\n    phase_timings[\"accelerator_init_sec\"] = float(time.perf_counter() - accelerator_init_t0)\n    loader_threads = max(1, int(args.num_workers))\n    eval_batch_size = int(args.eval_batch_size) if int(args.eval_batch_size) > 0 else int(train_cfg[\"batch_size\"])\n    device_transfer_mode = str(args.device_transfer_mode).strip().lower()\n    device_transfer_prefetch_batches = max(1, int(args.device_transfer_prefetch_batches))\n    # Keep host tensors pinned even on the threaded transfer path so H2D copies\n    # can overlap with compute on the prefetch stream.\n    use_pinned_transfer = bool(train_cfg.get(\"pin_memory\", False))\n    if int(args.loader_prefetch_batches) > 0:\n        loader_prefetch = max(1, int(args.loader_prefetch_batches))\n    else:\n        loader_prefetch = max(1, int(train_cfg.get(\"prefetch_factor\", 2)))\n    train_batch_overlap = max(0, int(args.train_batch_overlap))\n    skip_train_eval = bool(args.skip_train_eval)\n    skip_test = bool(args.skip_test)\n    print(\n        f\"[device-transfer] mode={device_transfer_mode} loader_pin_memory={int(use_pinned_transfer)} \"\n        f\"prefetch_batches={device_transfer_prefetch_batches}\",\n        flush=True,\n    )\n    if train_batch_overlap > 0:\n        print(\n            f\"[train-loader] batch_overlap={train_batch_overlap} batch_stride={int(train_cfg['batch_size']) - train_batch_overlap}\",\n            flush=True,\n        )\n    end_index_cache_dir = ensure_dir(Path(cfg[\"paths\"][\"output_root\"]) / \"training_seq\" / \"_end_index_cache\")\n    loader_init_t0 = time.perf_counter()\n    train_loader = SeqMemmapBatchLoader(\n        train_days,\n        seq_len=args.seq_len,\n        sample_stride=args.sample_stride,\n        batch_size=int(train_cfg[\"batch_size\"]),\n        shuffle=True,\n        max_samples=args.max_samples,\n        prefer_fp16=bool(args.prefer_fp16),\n        rank=accelerator.process_index,\n        world_size=accelerator.num_processes,\n        seed=args.seed,\n        pin_memory=use_pinned_transfer,\n        loader_threads=loader_threads,\n        prefetch_batches=loader_prefetch,\n        pad_last_batch=(accelerator.num_processes > 1),\n        batch_overlap=train_batch_overlap,\n        index_cache_dir=end_index_cache_dir,\n    )\n    phase_timings[\"train_loader_init_sec\"] = float(time.perf_counter() - loader_init_t0)\n    loader_init_t0 = time.perf_counter()\n    train_eval_loader = (\n        SeqMemmapBatchLoader(\n            train_days,\n            seq_len=args.seq_len,\n            sample_stride=args.sample_stride,\n            batch_size=eval_batch_size,\n            shuffle=False,\n            max_samples=args.max_samples,\n            prefer_fp16=bool(args.prefer_fp16),\n            rank=accelerator.process_index,\n            world_size=accelerator.num_processes,\n            seed=args.seed,\n            pin_memory=use_pinned_transfer,\n            loader_threads=loader_threads,\n            prefetch_batches=loader_prefetch,\n            pad_last_batch=False,\n            batch_overlap=0,\n            index_cache_dir=end_index_cache_dir,\n        )\n        if not profile_only and not skip_train_eval and train_batch_overlap > 0\n        else None\n    )\n    phase_timings[\"train_eval_loader_init_sec\"] = float(time.perf_counter() - loader_init_t0)\n    loader_init_t0 = time.perf_counter()\n    valid_loader = (\n        SeqMemmapBatchLoader(\n            valid_days,\n            seq_len=args.seq_len,\n            sample_stride=args.sample_stride,\n            batch_size=eval_batch_size,\n            shuffle=False,\n            max_samples=args.max_samples,\n            prefer_fp16=bool(args.prefer_fp16),\n            rank=accelerator.process_index,\n            world_size=accelerator.num_processes,\n            seed=args.seed,\n            pin_memory=use_pinned_transfer,\n            loader_threads=loader_threads,\n            prefetch_batches=loader_prefetch,\n            pad_last_batch=False,\n            batch_overlap=0,\n            index_cache_dir=end_index_cache_dir,\n        )\n        if not profile_only\n        else None\n    )\n    phase_timings[\"valid_loader_init_sec\"] = float(time.perf_counter() - loader_init_t0)\n    loader_init_t0 = time.perf_counter()\n    test_loader = (\n        SeqMemmapBatchLoader(\n            test_days,\n            seq_len=args.seq_len,\n            sample_stride=args.sample_stride,\n            batch_size=eval_batch_size,\n            shuffle=False,\n            max_samples=-1,\n            prefer_fp16=bool(args.prefer_fp16),\n            rank=accelerator.process_index,\n            world_size=accelerator.num_processes,\n            seed=args.seed,\n            pin_memory=use_pinned_transfer,\n            loader_threads=loader_threads,\n            prefetch_batches=loader_prefetch,\n            pad_last_batch=False,\n            batch_overlap=0,\n            index_cache_dir=end_index_cache_dir,\n        )\n        if test_days and not skip_test and not profile_only\n        else None\n    )\n    phase_timings[\"test_loader_init_sec\"] = float(time.perf_counter() - loader_init_t0)\n    stats_stride = max(20, args.sample_stride)\n    feature_stats_cache = build_feature_stats_cache_path(\n        output_root=Path(cfg[\"paths\"][\"output_root\"]),\n        row_root=row_root,\n        train_days=train_days,\n        prefer_fp16=bool(args.prefer_fp16),\n        sample_stride=stats_stride,\n    )\n    feature_stats_t0 = time.perf_counter()\n    if accelerator.is_main_process:\n        if feature_stats_cache.exists() and feature_stats_cache.stat().st_size > 0:\n            print(f\"[feature-stats] reused_from={feature_stats_cache}\", flush=True)\n            feat_mean_np, feat_std_np = load_feature_stats_cache(feature_stats_cache)\n        else:\n            print(\n                f\"[feature-stats] computing cache={feature_stats_cache} sample_stride={stats_stride}\",\n                flush=True,\n            )\n            feat_mean_np, feat_std_np = compute_feature_stats(\n                train_days, prefer_fp16=bool(args.prefer_fp16), sample_stride=stats_stride\n            )\n            tmp_path = feature_stats_cache.with_suffix(f\".tmp.{int(time.time())}.npz\")\n            np.savez(tmp_path, mean=feat_mean_np, std=feat_std_np)\n            tmp_path.replace(feature_stats_cache)\n            print(f\"[feature-stats] saved_cache={feature_stats_cache}\", flush=True)\n    else:\n        print(f\"[feature-stats] waiting_for_cache={feature_stats_cache}\", flush=True)\n        feat_mean_np, feat_std_np = wait_for_feature_stats_cache(feature_stats_cache)\n        print(f\"[feature-stats] loaded_cache={feature_stats_cache}\", flush=True)\n    accelerator.wait_for_everyone()\n    phase_timings[\"feature_stats_sec\"] = float(time.perf_counter() - feature_stats_t0)\n\n    model_init_t0 = time.perf_counter()\n    model = GRURegressor(\n        input_dim=200,\n        hidden_dim=args.hidden_dim,\n        num_layers=args.num_layers,\n        dropout=args.dropout,\n        pooling=args.pooling,\n        bidirectional=bool(args.bidirectional),\n        use_cnn1d=bool(getattr(args, \"use_cnn1d\", 0)),\n        input_gate_hidden_dim=int(getattr(args, \"input_gate_hidden_dim\", 0)),\n        input_gate_bias=float(getattr(args, \"input_gate_bias\", 2.0)),\n    )\n    optimizer = torch.optim.AdamW(\n        model.parameters(), lr=float(train_cfg[\"learning_rate\"]), weight_decay=float(args.weight_decay)\n    )\n    phase_timings[\"model_optimizer_init_sec\"] = float(time.perf_counter() - model_init_t0)\n\n    prepare_t0 = time.perf_counter()\n    if test_loader is not None:\n        model, optimizer = accelerator.prepare(model, optimizer)\n    else:\n        model, optimizer = accelerator.prepare(model, optimizer)\n    phase_timings[\"accelerator_prepare_sec\"] = float(time.perf_counter() - prepare_t0)\n    tensor_init_t0 = time.perf_counter()\n    feat_mean = torch.from_numpy(feat_mean_np).to(accelerator.device)[None, None, :]\n    feat_std = torch.from_numpy(feat_std_np).to(accelerator.device)[None, None, :]\n    seq_offsets = torch.arange(-(args.seq_len - 1), 1, dtype=torch.long, device=accelerator.device)\n    phase_timings[\"device_tensor_init_sec\"] = float(time.perf_counter() - tensor_init_t0)\n    device_prefetch_t0 = time.perf_counter()\n    if device_transfer_mode == \"thread_prefetch\":\n        train_loader = DeviceTransferPrefetchLoader(\n            train_loader,\n            device=accelerator.device,\n            prefetch_batches=device_transfer_prefetch_batches,\n            copy_non_blocking=use_pinned_transfer,\n        )\n        if train_eval_loader is not None:\n            train_eval_loader = DeviceTransferPrefetchLoader(\n                train_eval_loader,\n                device=accelerator.device,\n                prefetch_batches=device_transfer_prefetch_batches,\n                copy_non_blocking=use_pinned_transfer,\n            )\n        if valid_loader is not None:\n            valid_loader = DeviceTransferPrefetchLoader(\n                valid_loader,\n                device=accelerator.device,\n                prefetch_batches=device_transfer_prefetch_batches,\n                copy_non_blocking=use_pinned_transfer,\n            )\n        if test_loader is not None:\n            test_loader = DeviceTransferPrefetchLoader(\n                test_loader,\n                device=accelerator.device,\n                prefetch_batches=device_transfer_prefetch_batches,\n                copy_non_blocking=use_pinned_transfer,\n            )\n    elif device_transfer_mode == \"materialize_prefetch\":\n        train_loader = MaterializeDevicePrefetchLoader(\n            train_loader,\n            device=accelerator.device,\n            seq_offsets=seq_offsets,\n            feat_mean=feat_mean,\n            feat_std=feat_std,\n            prefetch_batches=device_transfer_prefetch_batches,\n            copy_non_blocking=use_pinned_transfer,\n        )\n        if train_eval_loader is not None:\n            train_eval_loader = MaterializeDevicePrefetchLoader(\n                train_eval_loader,\n                device=accelerator.device,\n                seq_offsets=seq_offsets,\n                feat_mean=feat_mean,\n                feat_std=feat_std,\n                prefetch_batches=device_transfer_prefetch_batches,\n                copy_non_blocking=use_pinned_transfer,\n            )\n        if valid_loader is not None:\n            valid_loader = MaterializeDevicePrefetchLoader(\n                valid_loader,\n                device=accelerator.device,\n                seq_offsets=seq_offsets,\n                feat_mean=feat_mean,\n                feat_std=feat_std,\n                prefetch_batches=device_transfer_prefetch_batches,\n                copy_non_blocking=use_pinned_transfer,\n            )\n        if test_loader is not None:\n            test_loader = MaterializeDevicePrefetchLoader(\n                test_loader,\n                device=accelerator.device,\n                seq_offsets=seq_offsets,\n                feat_mean=feat_mean,\n                feat_std=feat_std,\n                prefetch_batches=device_transfer_prefetch_batches,\n                copy_non_blocking=use_pinned_transfer,\n            )\n    phase_timings[\"device_prefetch_wrap_sec\"] = float(time.perf_counter() - device_prefetch_t0)\n\n    output_init_t0 = time.perf_counter()\n    out_root = ensure_dir(Path(cfg[\"paths\"][\"output_root\"]) / \"training_seq\" / args.run_name)\n    log_path = out_root / \"train_log.csv\"\n    best_model_path = out_root / \"gru_seq_memmap_best_val_ic.pt\"\n    save_topk_val_checkpoints = max(\n        1,\n        int(getattr(args, \"save_topk_val_checkpoints\", 1)),\n        int(getattr(args, \"test_checkpoint_rank\", 1)),\n    )\n    selected_test_checkpoint_rank = max(1, int(getattr(args, \"test_checkpoint_rank\", 1)))\n    if accelerator.is_main_process and not profile_only:\n        with log_path.open(\"w\", newline=\"\", encoding=\"utf-8\") as f:\n            csv.writer(f).writerow(\n                [\n                    \"epoch\",\n                    \"lr\",\n                    \"train_loss\",\n                    \"train_ic\",\n                    \"train_rmse\",\n                    \"val_loss\",\n                    \"val_ic\",\n                    \"val_rmse\",\n                    \"val_mae\",\n                    \"is_best\",\n                    \"pooling\",\n                ]\n            )\n    phase_timings[\"output_init_sec\"] = float(time.perf_counter() - output_init_t0)\n\n    epochs = int(train_cfg[\"epochs\"])\n    epoch_lrs = build_epoch_lr_schedule(train_cfg, epochs=epochs)\n    total_train_batches = 0\n    try:\n        total_train_batches = len(train_loader)\n    except TypeError:\n        total_train_batches = 0\n    train_progress_step = (\n        max(1, total_train_batches // max(1, int(args.train_progress_splits))) if total_train_batches > 0 else 0\n    )\n    best_epoch = 0\n    best_val_m: Dict[str, float] | None = None\n    final_val_m: Dict[str, float] | None = None\n    top_val_checkpoints: List[dict] = []\n    first_train_step_local: np.ndarray | None = None\n    collective_warmup_sec = 0.0\n    collective_warmup_done = False\n    train_eval_timing: dict | None = None\n    valid_eval_timing: dict | None = None\n    test_eval_timing: dict | None = None\n    if profile_only:\n        train_loader.set_epoch(1)\n        profile_summary = profile_train_steps(\n            model=model,\n            optimizer=optimizer,\n            loader=train_loader,\n            accelerator=accelerator,\n            feat_mean=feat_mean,\n            feat_std=feat_std,\n            seq_offsets=seq_offsets,\n            copy_non_blocking=use_pinned_transfer,\n            warmup_steps=profile_warmup_steps,\n            profile_steps=profile_steps,\n            grad_accum_steps=grad_accum_steps,\n        )\n        accelerator.wait_for_everyone()\n        if accelerator.is_main_process:\n            summary_path = out_root / \"step_profile_summary.json\"\n            summary = {\n                \"run_name\": args.run_name,\n                \"row_root\": str(row_root),\n                \"world_size\": int(accelerator.num_processes),\n                \"batch_size_per_rank\": int(train_cfg[\"batch_size\"]),\n                \"train_batch_overlap\": train_batch_overlap,\n                \"runtime_backend\": args.runtime_backend,\n                \"device_transfer_mode\": device_transfer_mode,\n                \"seq_len\": int(args.seq_len),\n                \"sample_stride\": int(args.sample_stride),\n                \"profile\": profile_summary,\n            }\n            summary_path.write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding=\"utf-8\")\n            print(\n                f\"[profile-train] saved={summary_path} global_step_sec={profile_summary['global_step_sec_estimate']:.6f} \"\n                f\"global_samples_per_sec={profile_summary['global_samples_per_sec_estimate']:.2f}\",\n                flush=True,\n            )\n        accelerator.close()\n        return\n    train_loop_t0 = time.perf_counter()\n    for epoch in range(1, epochs + 1):\n        current_lr = float(epoch_lrs[epoch - 1])\n        set_optimizer_lr(optimizer, current_lr)\n        epoch_train_t0 = time.perf_counter()\n        model.train()\n        train_loader.set_epoch(epoch)\n        if accelerator.is_main_process:\n            print(f\"[epoch {epoch}] lr={current_lr:.8f}\", flush=True)\n        train_iter = iter(train_loader)\n        if not collective_warmup_done:\n            collective_warmup_sec = warmup_collective(accelerator)\n            collective_warmup_done = True\n            phase_timings[\"collective_warmup_sec\"] = float(collective_warmup_sec)\n        batch_i = 0\n        accum_micro = 0\n        optimizer.zero_grad(set_to_none=True)\n        while True:\n            fetch_t0 = time.perf_counter()\n            try:\n                batch = next(train_iter)\n            except StopIteration:\n                break\n            fetch_t1 = time.perf_counter()\n            batch_i += 1\n            accum_micro += 1\n            last_batch = total_train_batches > 0 and batch_i == total_train_batches\n            should_step = accum_micro >= grad_accum_steps or last_batch\n            group_total = min(grad_accum_steps, accum_micro + max(0, total_train_batches - batch_i))\n            measure_first_step = bool(profile_phases and first_train_step_local is None and epoch == 1)\n            xb, yb, wb = resolve_model_batch(\n                batch,\n                seq_offsets=seq_offsets,\n                feat_mean=feat_mean,\n                feat_std=feat_std,\n                device=accelerator.device,\n                copy_non_blocking=use_pinned_transfer,\n            )\n            if measure_first_step:\n                sync_if_needed(accelerator.device)\n                mat_t = time.perf_counter()\n            sync_ctx = model.no_sync() if hasattr(model, \"no_sync\") and not should_step else nullcontext()\n            with sync_ctx:\n                with accelerator.autocast():\n                    pred = model(xb)\n                    loss_raw = weighted_mse(pred, yb, wb)\n                    loss = loss_raw / float(group_total)\n                if measure_first_step:\n                    sync_if_needed(accelerator.device)\n                    fwd_t = time.perf_counter()\n                accelerator.backward(loss)\n                if measure_first_step:\n                    sync_if_needed(accelerator.device)\n                    bwd_t = time.perf_counter()\n            if should_step:\n                accelerator.step_optimizer(optimizer)\n                optimizer.zero_grad(set_to_none=True)\n                accum_micro = 0\n            if measure_first_step:\n                sync_if_needed(accelerator.device)\n                step_t = time.perf_counter()\n                first_train_step_local = np.asarray(\n                    [\n                        fetch_t1 - fetch_t0,\n                        mat_t - fetch_t1,\n                        fwd_t - mat_t,\n                        bwd_t - fwd_t,\n                        step_t - bwd_t,\n                        step_t - fetch_t0,\n                    ],\n                    dtype=np.float64,\n                )\n            if (\n                accelerator.is_main_process\n                and total_train_batches > 0\n                and (batch_i % train_progress_step == 0 or batch_i == total_train_batches)\n            ):\n                print(\n                    f\"[epoch {epoch}] train progress {batch_i}/{total_train_batches} \"\n                    f\"loss={loss_raw.detach().float().item():.6f}\",\n                    flush=True,\n                )\n        phase_timings[f\"epoch_{epoch}_train_sec\"] = float(time.perf_counter() - epoch_train_t0)\n\n        if skip_train_eval:\n            train_m = {\n                \"loss\": float(\"nan\"),\n                \"ic\": float(\"nan\"),\n                \"rmse\": float(\"nan\"),\n                \"mae\": float(\"nan\"),\n                \"n\": 0,\n            }\n            if accelerator.is_main_process:\n                print(f\"[epoch {epoch}] train-eval skipped\", flush=True)\n        else:\n            train_eval_t0 = time.perf_counter()\n            train_m = evaluate(\n                model,\n                train_eval_loader if train_eval_loader is not None else train_loader,\n                feat_mean,\n                feat_std,\n                seq_offsets,\n                accelerator,\n                split_name=f\"epoch {epoch} train-eval\",\n                progress_parts=args.eval_progress_splits,\n                copy_non_blocking=use_pinned_transfer,\n                collect_timing=profile_phases,\n            )\n            phase_timings[f\"epoch_{epoch}_train_eval_sec\"] = float(time.perf_counter() - train_eval_t0)\n            train_eval_timing = train_m.pop(\"_timing\", None)\n        valid_eval_t0 = time.perf_counter()\n        val_m = evaluate(\n            model,\n            valid_loader,\n            feat_mean,\n            feat_std,\n            seq_offsets,\n            accelerator,\n            split_name=f\"epoch {epoch} valid\",\n            progress_parts=args.eval_progress_splits,\n            copy_non_blocking=use_pinned_transfer,\n            collect_timing=profile_phases,\n        )\n        phase_timings[f\"epoch_{epoch}_valid_eval_sec\"] = float(time.perf_counter() - valid_eval_t0)\n        valid_eval_timing = val_m.pop(\"_timing\", None)\n        final_val_m = dict(val_m)\n        top_val_checkpoints, entered_topk, dropped_topk = refresh_top_val_checkpoints(\n            top_val_checkpoints=top_val_checkpoints,\n            candidate_epoch=epoch,\n            candidate_val_metrics=val_m,\n            keep_topk=save_topk_val_checkpoints,\n            out_root=out_root,\n        )\n        is_best = 0\n        if best_val_m is None or float(val_m[\"ic\"]) > float(best_val_m[\"ic\"]):\n            best_val_m = dict(val_m)\n            best_epoch = epoch\n            is_best = 1\n            if accelerator.is_main_process:\n                torch.save(accelerator.unwrap_model(model).state_dict(), best_model_path)\n\n        if accelerator.is_main_process:\n            if entered_topk:\n                torch.save(\n                    accelerator.unwrap_model(model).state_dict(),\n                    build_val_epoch_checkpoint_path(out_root, epoch),\n                )\n            for dropped_item in dropped_topk:\n                dropped_path = Path(str(dropped_item[\"path\"]))\n                if dropped_path.exists():\n                    dropped_path.unlink()\n            with log_path.open(\"a\", newline=\"\", encoding=\"utf-8\") as f:\n                csv.writer(f).writerow(\n                    [\n                        epoch,\n                        current_lr,\n                        train_m[\"loss\"],\n                        train_m[\"ic\"],\n                        train_m[\"rmse\"],\n                        val_m[\"loss\"],\n                        val_m[\"ic\"],\n                        val_m[\"rmse\"],\n                        val_m[\"mae\"],\n                        is_best,\n                        str(args.pooling),\n                    ]\n                )\n            print(\n                f\"[epoch {epoch}] \"\n                f\"lr={current_lr:.8f} \"\n                f\"train_loss={train_m['loss']:.6f} train_ic={train_m['ic']:.6f} train_rmse={train_m['rmse']:.6f} \"\n                f\"val_loss={val_m['loss']:.6f} val_ic={val_m['ic']:.6f} val_rmse={val_m['rmse']:.6f} \"\n                f\"best_epoch={best_epoch}\"\n            )\n    phase_timings[\"fit_loop_sec\"] = float(time.perf_counter() - train_loop_t0)\n    if profile_phases and first_train_step_local is not None:\n        phase_details[\"first_train_step\"] = summarize_timing_across_ranks(\n            accelerator=accelerator,\n            labels=[\n                \"batch_wait_sec\",\n                \"materialize_sec\",\n                \"forward_sec\",\n                \"backward_sec\",\n                \"optimizer_step_sec\",\n                \"step_total_sec\",\n            ],\n            count=1,\n            sums=first_train_step_local,\n            count_key=\"measured_steps\",\n            global_batch=int(train_cfg[\"batch_size\"]) * int(accelerator.num_processes),\n        )\n    if collective_warmup_done:\n        phase_details[\"collective_warmup_sec\"] = {\"local_sec\": float(collective_warmup_sec)}\n    if train_eval_timing is not None:\n        phase_details[\"train_eval_timing\"] = train_eval_timing\n    if valid_eval_timing is not None:\n        phase_details[\"valid_eval_timing\"] = valid_eval_timing\n\n    def evaluate_saved_checkpoint(checkpoint_entry: dict, split_name: str) -> Tuple[Dict[str, float], float, dict | None]:\n        eval_t0 = time.perf_counter()\n        state_dict = torch.load(str(checkpoint_entry[\"path\"]), map_location=\"cpu\", weights_only=True)\n        accelerator.unwrap_model(model).load_state_dict(state_dict)\n        accelerator.wait_for_everyone()\n        if accelerator.is_main_process:\n            print(\n                f\"[{split_name}] evaluating val_rank={checkpoint_entry['rank']} \"\n                f\"epoch={checkpoint_entry['epoch']} val_ic={checkpoint_entry['val_ic']:.6f}\",\n                flush=True,\n            )\n        metrics = evaluate(\n            model,\n            test_loader,\n            feat_mean,\n            feat_std,\n            seq_offsets,\n            accelerator,\n            split_name=split_name,\n            progress_parts=args.eval_progress_splits,\n            copy_non_blocking=use_pinned_transfer,\n            collect_timing=profile_phases,\n        )\n        elapsed = float(time.perf_counter() - eval_t0)\n        timing = metrics.pop(\"_timing\", None)\n        return metrics, elapsed, timing\n\n    accelerator.wait_for_everyone()\n    test_m_best: Dict[str, float] | None = None\n    test_m_selected: Dict[str, float] | None = None\n    selected_test_checkpoint: dict | None = None\n    if test_loader is not None and not skip_test:\n        if not top_val_checkpoints:\n            raise RuntimeError(\"No validation checkpoints were recorded for test evaluation.\")\n        if selected_test_checkpoint_rank > len(top_val_checkpoints):\n            raise RuntimeError(\n                f\"Requested test_checkpoint_rank={selected_test_checkpoint_rank} but only \"\n                f\"{len(top_val_checkpoints)} validation checkpoints are available.\"\n            )\n        best_test_checkpoint = top_val_checkpoints[0]\n        selected_test_checkpoint = top_val_checkpoints[selected_test_checkpoint_rank - 1]\n        test_m_best, test_best_sec, test_eval_timing = evaluate_saved_checkpoint(\n            best_test_checkpoint,\n            split_name=\"test best-val\",\n        )\n        phase_timings[\"test_eval_best_val_sec\"] = float(test_best_sec)\n        if selected_test_checkpoint_rank == 1:\n            test_m_selected = dict(test_m_best)\n            phase_timings[\"test_eval_selected_val_rank_sec\"] = float(test_best_sec)\n            phase_timings[\"test_eval_sec\"] = float(test_best_sec)\n        else:\n            test_m_selected, test_selected_sec, test_eval_timing_selected = evaluate_saved_checkpoint(\n                selected_test_checkpoint,\n                split_name=f\"test val-rank-{selected_test_checkpoint_rank}\",\n            )\n            phase_timings[\"test_eval_selected_val_rank_sec\"] = float(test_selected_sec)\n            phase_timings[\"test_eval_sec\"] = float(test_best_sec + test_selected_sec)\n            if test_eval_timing_selected is not None:\n                phase_details[\"test_eval_timing_selected_val_rank\"] = test_eval_timing_selected\n        if accelerator.is_main_process:\n            print(\n                f\"[test best-val] rmse={test_m_best['rmse']:.6f} mae={test_m_best['mae']:.6f} ic={test_m_best['ic']:.6f}\",\n                flush=True,\n            )\n            if selected_test_checkpoint_rank != 1 and test_m_selected is not None:\n                print(\n                    f\"[test val-rank-{selected_test_checkpoint_rank}] \"\n                    f\"rmse={test_m_selected['rmse']:.6f} mae={test_m_selected['mae']:.6f} ic={test_m_selected['ic']:.6f}\",\n                    flush=True,\n                )\n    elif test_loader is not None and accelerator.is_main_process:\n        print(\"[test] skipped\", flush=True)\n    if test_eval_timing is not None:\n        phase_details[\"test_eval_timing\"] = test_eval_timing\n    final_save_t0 = time.perf_counter()\n    if accelerator.is_main_process:\n        torch.save(accelerator.unwrap_model(model).state_dict(), out_root / \"gru_seq_memmap_ddp.pt\")\n    accelerator.wait_for_everyone()\n    phase_timings[\"final_save_sec\"] = float(time.perf_counter() - final_save_t0)\n    phase_timings[\"total_run_sec\"] = float(time.perf_counter() - run_t0)\n    if accelerator.is_main_process:\n        summary = {\n            \"run_name\": args.run_name,\n            \"cache_name\": args.cache_name,\n            \"max_days\": args.max_days,\n            \"train_pool_start_date\": args.train_pool_start_date,\n            \"train_pool_end_date\": args.train_pool_end_date,\n            \"test_start_date\": args.test_start_date,\n            \"test_end_date\": args.test_end_date,\n            \"valid_split_ratio\": split_ratio,\n            \"selection_metric\": \"val_ic\",\n            \"train_day_count\": len(train_days),\n            \"valid_day_count\": len(valid_days),\n            \"test_day_count\": len(test_days),\n            \"train_samples\": int(train_loader.total_samples),\n            \"valid_samples\": int(valid_loader.total_samples),\n            \"test_samples\": int(test_loader.total_samples) if test_loader is not None else 0,\n            \"seed\": int(args.seed),\n            \"seq_len\": args.seq_len,\n            \"sample_stride\": args.sample_stride,\n            \"hidden_dim\": args.hidden_dim,\n            \"num_layers\": args.num_layers,\n            \"bidirectional\": bool(args.bidirectional),\n            \"use_cnn1d\": bool(getattr(args, \"use_cnn1d\", 0)),\n            \"input_gate_hidden_dim\": int(getattr(args, \"input_gate_hidden_dim\", 0)),\n            \"input_gate_bias\": float(getattr(args, \"input_gate_bias\", 2.0)),\n            \"weight_decay\": args.weight_decay,\n            \"prefer_fp16\": bool(args.prefer_fp16),\n            \"use_amp\": bool(args.use_amp),\n            \"pooling\": str(args.pooling),\n            \"loader_backend\": \"threaded_seq_batch_loader\",\n            \"loader_threads\": loader_threads,\n            \"loader_prefetch_batches\": loader_prefetch,\n            \"train_batch_overlap\": train_batch_overlap,\n            \"train_batch_stride\": int(train_cfg[\"batch_size\"]) - train_batch_overlap,\n            \"eval_batch_size_per_rank\": eval_batch_size,\n            \"grad_accum_steps\": grad_accum_steps,\n            \"effective_global_batch\": int(train_cfg[\"batch_size\"]) * int(accelerator.num_processes) * grad_accum_steps,\n            \"runtime_backend\": args.runtime_backend,\n            \"device_transfer_mode\": device_transfer_mode,\n            \"device_transfer_prefetch_batches\": device_transfer_prefetch_batches,\n            \"pinned_nonblocking_transfer\": use_pinned_transfer,\n            \"skip_train_eval\": skip_train_eval,\n            \"skip_test\": skip_test,\n            \"train_cfg\": train_cfg,\n            \"epoch_learning_rates\": [float(x) for x in epoch_lrs],\n            \"loss\": \"weighted_mse\",\n            \"metric_main\": \"val_ic\",\n            \"test_selection_rule\": f\"val_ic_rank_{selected_test_checkpoint_rank}\",\n            \"save_topk_val_checkpoints\": int(save_topk_val_checkpoints),\n            \"best_epoch_by_val_ic\": best_epoch,\n            \"best_val_metrics\": best_val_m,\n            \"final_epoch_val_metrics\": final_val_m,\n            \"top_val_checkpoints\": top_val_checkpoints,\n            \"selected_test_checkpoint_rank\": int(selected_test_checkpoint_rank),\n            \"selected_test_checkpoint_epoch\": int(selected_test_checkpoint[\"epoch\"]) if selected_test_checkpoint is not None else None,\n            \"selected_test_model_path\": str(selected_test_checkpoint[\"path\"]) if selected_test_checkpoint is not None else None,\n            \"test_metrics_at_best_val\": test_m_best,\n            \"test_metrics_at_selected_val_rank\": test_m_selected,\n            \"best_model_path\": str(best_model_path),\n            \"phase_timings\": phase_timings,\n            \"phase_details\": phase_details,\n        }\n        (out_root / \"training_summary.json\").write_text(json.dumps(summary, ensure_ascii=False, indent=2), encoding=\"utf-8\")\n        np.savez(out_root / \"feature_stats.npz\", mean=feat_mean_np, std=feat_std_np)\n        if profile_phases:\n            phase_summary = {\n                \"run_name\": args.run_name,\n                \"row_root\": str(row_root),\n                \"world_size\": int(accelerator.num_processes),\n                \"batch_size_per_rank\": int(train_cfg[\"batch_size\"]),\n                \"train_batch_overlap\": train_batch_overlap,\n                \"runtime_backend\": args.runtime_backend,\n                \"device_transfer_mode\": device_transfer_mode,\n                \"phase_timings\": phase_timings,\n                \"phase_details\": phase_details,\n            }\n            phase_path = out_root / \"phase_profile_summary.json\"\n            phase_path.write_text(json.dumps(phase_summary, ensure_ascii=False, indent=2), encoding=\"utf-8\")\n            print(f\"[phase-profile] saved={phase_path}\", flush=True)\n    accelerator.close()\n\n\ndef main():\n    p = argparse.ArgumentParser()\n    p.add_argument(\"--config\", default=\"/intern9/huhongkai/hs300_factor_lab/configs/experiment_2024_2025_memmap.json\")\n    p.add_argument(\"--run-name\", default=\"seq_gru_memmap_run\")\n    p.add_argument(\"--cache-name\", default=\"top200_eps20_rows_2024_2025\")\n    p.add_argument(\n        \"--row-root-override\",\n        type=str,\n        default=\"\",\n        help=\"Optional row_memmap root or exact cache dir. Useful for local staged cache under /dev/shm.\",\n    )\n    p.add_argument(\"--max-days\", type=int, default=-1)\n    p.add_argument(\"--max-samples\", type=int, default=-1)\n    p.add_argument(\"--train-day-limit\", type=int, default=-1)\n    p.add_argument(\"--valid-day-limit\", type=int, default=-1)\n    p.add_argument(\"--test-day-limit\", type=int, default=-1)\n    p.add_argument(\"--train-pool-start-date\", type=str, default=\"\")\n    p.add_argument(\"--train-pool-end-date\", type=str, default=\"\")\n    p.add_argument(\"--test-start-date\", type=str, default=\"\")\n    p.add_argument(\"--test-end-date\", type=str, default=\"\")\n    p.add_argument(\"--valid-split-ratio\", type=float, default=-1.0)\n    p.add_argument(\"--seed\", type=int, default=20260312)\n    p.add_argument(\"--seq-len\", type=int, default=60)\n    p.add_argument(\"--sample-stride\", type=int, default=10)\n    p.add_argument(\"--train-batch-overlap\", type=int, default=0)\n    p.add_argument(\"--eval-batch-size\", type=int, default=0)\n    p.add_argument(\"--grad-accum-steps\", type=int, default=1)\n    p.add_argument(\"--runtime-backend\", type=str, default=\"native\", choices=[\"accelerate\", \"native\"])\n    p.add_argument(\n        \"--device-transfer-mode\",\n        type=str,\n        default=\"thread_prefetch\",\n        choices=[\"direct\", \"thread_prefetch\", \"materialize_prefetch\"],\n    )\n    p.add_argument(\"--device-transfer-prefetch-batches\", type=int, default=2)\n    p.add_argument(\"--loader-prefetch-batches\", type=int, default=0)\n    p.add_argument(\"--prefer-fp16\", type=int, default=1)\n    p.add_argument(\"--override-epochs\", type=int, default=-1)\n    p.add_argument(\"--override-lr\", type=float, default=-1.0)\n    p.add_argument(\"--override-batch-size\", type=int, default=-1)\n    p.add_argument(\"--hidden-dim\", type=int, default=256)\n    p.add_argument(\"--num-layers\", type=int, default=2)\n    p.add_argument(\"--dropout\", type=float, default=0.1)\n    p.add_argument(\"--pooling\", type=str, default=\"last\", choices=[\"last\", \"attn\"])\n    p.add_argument(\"--bidirectional\", type=int, default=0, help=\"1 to enable bidirectional GRU\")\n    p.add_argument(\"--use-cnn1d\", type=int, default=0, help=\"1 to use parallel 1D-CNN before GRU\")\n    p.add_argument(\"--input-gate-hidden-dim\", type=int, default=0, help=\">0 to enable feature/channel gate before GRU\")\n    p.add_argument(\"--input-gate-bias\", type=float, default=2.0, help=\"Initial bias for feature/channel gate\")\n    p.add_argument(\"--weight-decay\", type=float, default=1e-5)\n    p.add_argument(\"--num-workers\", type=int, default=4)\n    p.add_argument(\"--use-amp\", type=int, default=1, help=\"1 to enable fp16 mixed precision\")\n    p.add_argument(\"--train-progress-splits\", type=int, default=10)\n    p.add_argument(\"--eval-progress-splits\", type=int, default=4)\n    p.add_argument(\"--skip-train-eval\", type=int, default=1, help=\"1 to skip full train-set evaluation\")\n    p.add_argument(\"--skip-test\", type=int, default=0, help=\"1 to skip final test evaluation\")\n    p.add_argument(\"--profile-phases\", type=int, default=0)\n    p.add_argument(\"--profile-warmup-steps\", type=int, default=0)\n    p.add_argument(\"--profile-steps\", type=int, default=0)\n    p.add_argument(\"--stop-after-profile\", type=int, default=0)\n    p.add_argument(\"--save-topk-val-checkpoints\", type=int, default=1)\n    p.add_argument(\"--test-checkpoint-rank\", type=int, default=1)\n    args = p.parse_args()\n    run(args)\n\n\nif __name__ == \"__main__\":\n    main()\n","message":"The file /intern9/huhongkai/hs300_factor_lab/src/train_seq_gru_ddp_memmap.py has been updated."}},"isError":false}}}}