{"role":"tool","content":[{"type":"tool-result","toolCallId":"toolu_01QZJsSMiRGcJ31ogGAEdkmD","toolName":"StrReplace","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":"toolu_01QZJsSMiRGcJ31ogGAEdkmD","providerOptions":{"cursor":{"highLevelToolCallResult":{"output":{"success":{"path":"/intern9/huhongkai/hs300_factor_lab/src/train_seq_gru_ddp_memmap.py","linesAdded":1,"linesRemoved":1,"diffString":"                     bwd_t = time.perf_counter()\n             if should_step:\n-                accelerator.step_optimizer(optimizer)\n+                accelerator.step_optimizer(optimizer, model=model, grad_clip_norm=grad_clip_norm)\n                 optimizer.zero_grad(set_to_none=True)\n                 accum_micro = 0","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\nfrom splitview_time_utils import load_split_view_source_field_slice\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, lr_schedule: str = \"constant\", lr_warmup_epochs: int = 0) -> List[float]:\n    import math\n    base_lr = float(train_cfg[\"learning_rate\"])\n    raw_values = train_cfg.get(\"lr_epoch_values\")\n    if raw_values is not None:\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    epochs = int(epochs)\n    warmup = min(max(0, int(lr_warmup_epochs)), epochs)\n    if lr_schedule == \"cosine\":\n        lrs = []\n        for e in range(epochs):\n            if e < warmup:\n                lrs.append(base_lr * (e + 1) / max(1, warmup))\n            else:\n                progress = (e - warmup) / max(1, epochs - warmup - 1)\n                lrs.append(base_lr * 0.5 * (1.0 + math.cos(math.pi * progress)))\n        return lrs\n    return [base_lr for _ in range(epochs)]\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\n_REGRESSION_LOSS_NAME = \"mse\"\n_HUBER_DELTA = 1.0\n\n\ndef configure_regression_loss(loss_name: str, huber_delta: float) -> None:\n    global _REGRESSION_LOSS_NAME, _HUBER_DELTA\n    normalized = str(loss_name).strip().lower()\n    if normalized not in {\"mse\", \"huber\"}:\n        raise ValueError(f\"Unsupported regression loss: {loss_name}\")\n    _REGRESSION_LOSS_NAME = normalized\n    _HUBER_DELTA = max(1e-8, float(huber_delta))\n\n\ndef regression_loss_name() -> str:\n    return str(_REGRESSION_LOSS_NAME)\n\n\ndef regression_huber_delta() -> float:\n    return float(_HUBER_DELTA)\n\n\ndef pointwise_regression_loss(pred: torch.Tensor, y: torch.Tensor) -> torch.Tensor:\n    if regression_loss_name() == \"huber\":\n        return torch.nn.functional.huber_loss(pred, y, reduction=\"none\", delta=regression_huber_delta())\n    return (pred - y) ** 2\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, model=None, grad_clip_norm: float = 0.0) -> None:\n        if grad_clip_norm > 0.0 and model is not None:\n            self.accelerator.clip_grad_norm_(model.parameters(), grad_clip_norm)\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, model=None, grad_clip_norm: float = 0.0) -> None:\n        if self.use_amp:\n            self.scaler.unscale_(optimizer)\n        if grad_clip_norm > 0.0 and model is not None:\n            torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip_norm)\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] = {}\nDEFAULT_MIN_TIMECODE = 93_000_000\n_USE_SOURCE_RAW_WEIGHTS = False\n\n\ndef configure_split_view_weight_loading(use_source_raw_weights: bool) -> None:\n    global _USE_SOURCE_RAW_WEIGHTS\n    _USE_SOURCE_RAW_WEIGHTS = bool(use_source_raw_weights)\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    if bool(_USE_SOURCE_RAW_WEIGHTS):\n        try:\n            w = np.asarray(load_split_view_source_field_slice(day_dir, \"w\"), dtype=np.float32)\n        except Exception:\n            w = np.memmap(day_dir / \"w_float32.memmap\", mode=\"r\", dtype=np.float32, shape=(n_rows,))\n    else:\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 filter_valid_end_indices(\n    day_dir: Path,\n    end_indices: np.ndarray,\n    min_timecode: int = -1,\n    require_positive_weight: bool = False,\n) -> np.ndarray:\n    ends = np.asarray(end_indices, dtype=np.int64)\n    if ends.size == 0:\n        return ends\n    mask = np.ones((int(ends.shape[0]),), dtype=bool)\n    if bool(require_positive_weight):\n        meta = load_day_meta(day_dir)\n        n_rows = int(meta[\"n_rows\"])\n        cache_w = np.memmap(day_dir / \"w_float32.memmap\", mode=\"r\", dtype=np.float32, shape=(n_rows,))\n        weight_mask = np.asarray(cache_w[ends], dtype=np.float32)\n        mask &= np.isfinite(weight_mask) & (weight_mask > 0.0)\n    if int(min_timecode) > 0:\n        dt_slice = np.asarray(load_split_view_source_field_slice(day_dir, \"datetime\"), dtype=np.int64)\n        dt_end = np.asarray(dt_slice[ends], dtype=np.int64)\n        mask &= dt_end >= int(min_timecode)\n    return ends[mask]\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(\n    cache_dir: Path,\n    day_dir: Path,\n    seq_len: int,\n    sample_stride: int,\n    min_timecode: int = -1,\n    require_positive_weight: bool = False,\n) -> 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            str(int(min_timecode)),\n            str(int(bool(require_positive_weight))),\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    min_timecode: int = -1,\n    require_positive_weight: bool = False,\n) -> np.ndarray:\n    if cache_dir is None:\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        return filter_valid_end_indices(\n            day_dir,\n            ends,\n            min_timecode=min_timecode,\n            require_positive_weight=require_positive_weight,\n        )\n    cache_path = build_end_index_cache_path(\n        cache_dir,\n        day_dir,\n        seq_len=seq_len,\n        sample_stride=sample_stride,\n        min_timecode=min_timecode,\n        require_positive_weight=require_positive_weight,\n    )\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    ends = filter_valid_end_indices(\n        day_dir,\n        ends,\n        min_timecode=min_timecode,\n        require_positive_weight=require_positive_weight,\n    )\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        min_timecode: int = -1,\n        require_positive_weight: bool = False,\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.min_timecode = int(min_timecode)\n        self.require_positive_weight = bool(require_positive_weight)\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                min_timecode=self.min_timecode,\n                require_positive_weight=self.require_positive_weight,\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        \n        # Optimize for sample_stride > 1: materialize sequences on CPU to reduce GPU transfer\n        cpu_materialize_threshold = 2.5\n        span_rows = span_end - span_start + 1\n        batch_rows = len(batch_end)\n        use_cpu_materialize = (\n            span_rows > int(batch_rows * self.seq_len * cpu_materialize_threshold)\n        )\n        \n        if use_cpu_materialize:\n            # CPU-side materialization: directly build (batch_size, seq_len, feat_dim) sequences\n            span_x_dtype = np.float32 if store.need_clean else store.x.dtype\n            batch_size = len(batch_end)\n            feat_dim = store.x.shape[1]\n            materialized_x = np.empty((batch_size, self.seq_len, feat_dim), dtype=span_x_dtype)\n            \n            for i, end_idx in enumerate(batch_end):\n                seq_start = int(end_idx) - self.seq_len + 1\n                seq_slice = store.x[seq_start : int(end_idx) + 1]\n                materialized_x[i, :, :] = seq_slice\n            \n            if store.need_clean:\n                np.nan_to_num(materialized_x, copy=False, nan=0.0, posinf=0.0, neginf=0.0)\n            \n            # For CPU-materialized case, span_x IS the materialized sequences, local_end is not needed for gather\n            span_xb = torch.from_numpy(np.ascontiguousarray(materialized_x))\n            local_endb = torch.empty((0,), dtype=torch.int64)  # sentinel: empty means already materialized\n        else:\n            # Original span-based approach for small strides\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            span_xb = torch.from_numpy(np.ascontiguousarray(span_x))\n            local_endb = torch.from_numpy(local_end)\n        \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        np.nan_to_num(y, copy=False, nan=0.0, posinf=0.0, neginf=0.0)\n        np.nan_to_num(w, copy=False, nan=0.0, posinf=0.0, neginf=0.0)\n        np.maximum(w, 0.0, out=w)\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            if local_endb.numel() > 0:\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        \n        # Handle CPU-materialized case differently\n        if part.local_end.numel() == 0:\n            # span_x is already (batch_size, seq_len, feat_dim), pad along batch dim\n            span_x = torch.cat([part.span_x, part.span_x[-1:].repeat(int(pad_size), 1, 1)], dim=0)\n            local_end = part.local_end  # keep empty sentinel\n        else:\n            # Original span-based: only pad local_end, y, w\n            span_x = part.span_x\n            local_end = torch.cat([part.local_end, part.local_end[-1:].repeat(int(pad_size))], dim=0)\n        \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        \n        if self.pin_memory:\n            span_x = span_x.pin_memory()\n            if local_end.numel() > 0:\n                local_end = local_end.pin_memory()\n            y = y.pin_memory()\n            w = w.pin_memory()\n        return PackedSeqBatchPart(span_x=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 * pointwise_regression_loss(pred, y)).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_unweighted_ic\": float(candidate_val_metrics.get(\"unweighted_ic\", candidate_val_metrics[\"ic\"])),\n        \"val_weighted_ic\": float(candidate_val_metrics.get(\"weighted_ic\", 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.get(\"val_weighted_ic\", 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        else:\n            ww = None\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        if ww is None:\n            weight = mask.to(dtype=torch.float64)\n        else:\n            weight = torch.where(mask, ww, zero).to(dtype=torch.float64)\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                weight.sum(),\n                (weight * p).sum(),\n                (weight * t).sum(),\n                (weight * p * p).sum(),\n                (weight * t * t).sum(),\n                (weight * p * t).sum(),\n                (weight * torch.abs(d)).sum(),\n                (weight * 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((16,), dtype=torch.float64, device=device)\n        return self.buf.to(device=device, dtype=torch.float64)\n\n\ndef corr_from_sums(sum_x: float, sum_y: float, sum_xx: float, sum_yy: float, sum_xy: float, denom_weight: float) -> float:\n    if (not np.isfinite(denom_weight)) or denom_weight <= 0:\n        return float(\"nan\")\n    mean_x = sum_x / denom_weight\n    mean_y = sum_y / denom_weight\n    var_x = max(sum_xx / denom_weight - mean_x * mean_x, 1e-12)\n    var_y = max(sum_yy / denom_weight - mean_y * mean_y, 1e-12)\n    cov_xy = sum_xy / denom_weight - mean_x * mean_y\n    return float(cov_xy / math.sqrt(var_x * var_y))\n\n\ndef metrics_from_tensor(t: torch.Tensor) -> dict:\n    values = [float(x) for x in t.tolist()]\n    if len(values) < 8:\n        raise ValueError(f\"metrics_from_tensor expects at least 8 values, got {len(values)}\")\n    n, sum_p, sum_y, sum_pp, sum_yy, sum_py, sum_abs, sum_sq = values[:8]\n    if (not np.isfinite(n)) or n <= 0:\n        return {\n            \"mse\": float(\"nan\"),\n            \"rmse\": float(\"nan\"),\n            \"mae\": float(\"nan\"),\n            \"ic\": float(\"nan\"),\n            \"unweighted_mse\": float(\"nan\"),\n            \"unweighted_rmse\": float(\"nan\"),\n            \"unweighted_mae\": float(\"nan\"),\n            \"unweighted_ic\": float(\"nan\"),\n            \"weighted_mse\": float(\"nan\"),\n            \"weighted_rmse\": float(\"nan\"),\n            \"weighted_mae\": float(\"nan\"),\n            \"weighted_ic\": float(\"nan\"),\n            \"weight_sum\": 0.0,\n            \"n\": 0,\n        }\n    mse = sum_sq / n\n    mae = sum_abs / n\n    ic = corr_from_sums(sum_p, sum_y, sum_pp, sum_yy, sum_py, n)\n    if len(values) >= 16:\n        weight_sum, sum_wp, sum_wy, sum_wpp, sum_wyy, sum_wpy, sum_wabs, sum_wsq = values[8:16]\n    elif len(values) >= 14:\n        weight_sum, sum_wp, sum_wy, sum_wpp, sum_wyy, sum_wpy = values[8:14]\n        sum_wabs, sum_wsq = sum_abs, sum_sq\n    else:\n        weight_sum, sum_wp, sum_wy, sum_wpp, sum_wyy, sum_wpy = n, sum_p, sum_y, sum_pp, sum_yy, sum_py\n        sum_wabs, sum_wsq = sum_abs, sum_sq\n    weighted_ic = corr_from_sums(sum_wp, sum_wy, sum_wpp, sum_wyy, sum_wpy, weight_sum)\n    weighted_mse = (sum_wsq / weight_sum) if weight_sum > 0 else float(\"nan\")\n    weighted_mae = (sum_wabs / weight_sum) if weight_sum > 0 else float(\"nan\")\n    weighted_rmse = math.sqrt(weighted_mse) if np.isfinite(weighted_mse) and weighted_mse >= 0 else float(\"nan\")\n    return {\n        \"mse\": float(weighted_mse),\n        \"rmse\": float(weighted_rmse),\n        \"mae\": float(weighted_mae),\n        \"ic\": float(weighted_ic),\n        \"unweighted_mse\": float(mse),\n        \"unweighted_rmse\": float(math.sqrt(mse)),\n        \"unweighted_mae\": float(mae),\n        \"unweighted_ic\": float(ic),\n        \"weighted_mse\": float(weighted_mse),\n        \"weighted_rmse\": float(weighted_rmse),\n        \"weighted_mae\": float(weighted_mae),\n        \"weighted_ic\": float(weighted_ic),\n        \"weight_sum\": float(weight_sum),\n        \"n\": int(n),\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    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            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    stats_t = accelerator.reduce(stats.to_tensor(accelerator.device), reduction=\"sum\")\n    m = metrics_from_tensor(stats_t)\n    m[\"loss\"] = float(m.get(\"weighted_mse\", float(\"nan\")))\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        yb = part.y.to(device, non_blocking=copy_non_blocking)\n        wb = part.w.to(device, non_blocking=copy_non_blocking)\n        \n        # Check if part is already CPU-materialized (local_end.numel() == 0 is sentinel)\n        if part.local_end.numel() == 0:\n            # Already materialized on CPU: span_x is (batch_size, seq_len, feat_dim)\n            xb = part.span_x.to(device, non_blocking=copy_non_blocking).float()\n            xb = (xb - feat_mean) / feat_std\n            xb = xb.contiguous()\n        else:\n            # Original span-based gather on GPU\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            seq_idx = local_end[:, None] + seq_offsets[None, :]\n            xb = span_x[seq_idx].float()\n            xb = (xb - feat_mean) / feat_std\n            xb = xb.contiguous()\n        \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    configure_regression_loss(getattr(args, \"regression_loss\", \"mse\"), getattr(args, \"huber_delta\", 1.0))\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    configure_split_view_weight_loading(bool(args.use_source_raw_weights))\n    min_timecode = int(args.min_timecode)\n    require_positive_weight = bool(args.require_positive_weight)\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    print(\n        f\"[sample-filter] min_timecode={min_timecode} require_positive_weight={int(require_positive_weight)} \"\n        f\"use_source_raw_weights={int(bool(args.use_source_raw_weights))}\",\n        flush=True,\n    )\n    print(\n        f\"[loss] regression_loss={regression_loss_name()} huber_delta={regression_huber_delta():.6f}\",\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        min_timecode=min_timecode,\n        require_positive_weight=require_positive_weight,\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            min_timecode=min_timecode,\n            require_positive_weight=require_positive_weight,\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            min_timecode=min_timecode,\n            require_positive_weight=require_positive_weight,\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            min_timecode=min_timecode,\n            require_positive_weight=require_positive_weight,\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    init_checkpoint = str(getattr(args, \"init_checkpoint\", \"\") or \"\").strip()\n    init_checkpoint_path: Path | None = None\n    if init_checkpoint:\n        init_checkpoint_path = Path(init_checkpoint).expanduser()\n        if not init_checkpoint_path.is_absolute():\n            init_checkpoint_path = (Path.cwd() / init_checkpoint_path).resolve()\n        if not init_checkpoint_path.is_file():\n            raise FileNotFoundError(f\"init checkpoint does not exist: {init_checkpoint_path}\")\n        state_dict = torch.load(str(init_checkpoint_path), map_location=\"cpu\", weights_only=True)\n        model.load_state_dict(state_dict)\n        if accelerator.is_main_process:\n            print(f\"[init-checkpoint] loaded={init_checkpoint_path}\", flush=True)\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_unweighted_ic\",\n                    \"train_weighted_ic\",\n                    \"train_rmse\",\n                    \"train_unweighted_rmse\",\n                    \"val_loss\",\n                    \"val_ic\",\n                    \"val_unweighted_ic\",\n                    \"val_weighted_ic\",\n                    \"val_rmse\",\n                    \"val_unweighted_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_start = max(1, int(getattr(args, \"epoch_start\", 1)))\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(epoch_start)\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                \"min_timecode\": int(min_timecode),\n                \"require_positive_weight\": bool(require_positive_weight),\n                \"use_source_raw_weights\": bool(args.use_source_raw_weights),\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 local_epoch in range(1, epochs + 1):\n        epoch = epoch_start + local_epoch - 1\n        current_lr = float(epoch_lrs[local_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 local_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                \"unweighted_ic\": float(\"nan\"),\n                \"weighted_ic\": float(\"nan\"),\n                \"rmse\": float(\"nan\"),\n                \"mae\": float(\"nan\"),\n                \"weight_sum\": 0.0,\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        current_weighted_ic = float(val_m.get(\"weighted_ic\", val_m[\"ic\"]))\n        best_weighted_ic = float(best_val_m.get(\"weighted_ic\", best_val_m[\"ic\"])) if best_val_m is not None else float(\"-inf\")\n        if best_val_m is None or current_weighted_ic > best_weighted_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.get(\"unweighted_ic\", float(\"nan\")),\n                        train_m.get(\"weighted_ic\", float(\"nan\")),\n                        train_m[\"rmse\"],\n                        train_m.get(\"unweighted_rmse\", float(\"nan\")),\n                        val_m[\"loss\"],\n                        val_m[\"ic\"],\n                        val_m.get(\"unweighted_ic\", float(\"nan\")),\n                        val_m.get(\"weighted_ic\", float(\"nan\")),\n                        val_m[\"rmse\"],\n                        val_m.get(\"unweighted_rmse\", float(\"nan\")),\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} \"\n                f\"train_uic={train_m.get('unweighted_ic', float('nan')):.6f} train_rmse={train_m['rmse']:.6f} \"\n                f\"val_loss={val_m['loss']:.6f} val_ic={val_m['ic']:.6f} \"\n                f\"val_uic={val_m.get('unweighted_ic', float('nan')):.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                f\"val_wic={checkpoint_entry.get('val_weighted_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} \"\n                f\"ic={test_m_best['ic']:.6f} uic={test_m_best.get('unweighted_ic', float('nan')):.6f} \"\n                f\"wic={test_m_best.get('weighted_ic', float('nan')):.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} \"\n                    f\"ic={test_m_selected['ic']:.6f} uic={test_m_selected.get('unweighted_ic', float('nan')):.6f} \"\n                    f\"wic={test_m_selected.get('weighted_ic', float('nan')):.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            \"epoch_start\": int(epoch_start),\n            \"epochs_this_run\": int(epochs),\n            \"init_checkpoint\": str(init_checkpoint_path) if init_checkpoint_path is not None else \"\",\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_weighted_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            \"min_timecode\": int(min_timecode),\n            \"require_positive_weight\": bool(require_positive_weight),\n            \"use_source_raw_weights\": bool(args.use_source_raw_weights),\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\": f\"weighted_{regression_loss_name()}\",\n            \"loss_config\": {\n                \"regression_loss\": regression_loss_name(),\n                \"huber_delta\": regression_huber_delta(),\n            },\n            \"metric_main\": \"val_weighted_ic\",\n            \"test_selection_rule\": f\"val_weighted_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_epoch_by_val_weighted_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                \"min_timecode\": int(min_timecode),\n                \"require_positive_weight\": bool(require_positive_weight),\n                \"use_source_raw_weights\": bool(args.use_source_raw_weights),\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(\"--regression-loss\", type=str, default=\"mse\", choices=[\"mse\", \"huber\"])\n    p.add_argument(\"--huber-delta\", type=float, default=1.0)\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(\n        \"--use-source-raw-weights\",\n        type=int,\n        default=1,\n        help=\"1 to load raw continuous weights from the source split npy instead of cache-side binary masks.\",\n    )\n    p.add_argument(\n        \"--min-timecode\",\n        type=int,\n        default=DEFAULT_MIN_TIMECODE,\n        help=\"Only keep sequence end points whose source datetime >= this HHMMSSmmm timecode.\",\n    )\n    p.add_argument(\n        \"--require-positive-weight\",\n        type=int,\n        default=1,\n        help=\"1 to drop sequence end points whose cached/source weight is not positive.\",\n    )\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    p.add_argument(\"--init-checkpoint\", type=str, default=\"\")\n    p.add_argument(\"--epoch-start\", 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\nfrom splitview_time_utils import load_split_view_source_field_slice\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, lr_schedule: str = \"constant\", lr_warmup_epochs: int = 0) -> List[float]:\n    import math\n    base_lr = float(train_cfg[\"learning_rate\"])\n    raw_values = train_cfg.get(\"lr_epoch_values\")\n    if raw_values is not None:\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    epochs = int(epochs)\n    warmup = min(max(0, int(lr_warmup_epochs)), epochs)\n    if lr_schedule == \"cosine\":\n        lrs = []\n        for e in range(epochs):\n            if e < warmup:\n                lrs.append(base_lr * (e + 1) / max(1, warmup))\n            else:\n                progress = (e - warmup) / max(1, epochs - warmup - 1)\n                lrs.append(base_lr * 0.5 * (1.0 + math.cos(math.pi * progress)))\n        return lrs\n    return [base_lr for _ in range(epochs)]\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\n_REGRESSION_LOSS_NAME = \"mse\"\n_HUBER_DELTA = 1.0\n\n\ndef configure_regression_loss(loss_name: str, huber_delta: float) -> None:\n    global _REGRESSION_LOSS_NAME, _HUBER_DELTA\n    normalized = str(loss_name).strip().lower()\n    if normalized not in {\"mse\", \"huber\"}:\n        raise ValueError(f\"Unsupported regression loss: {loss_name}\")\n    _REGRESSION_LOSS_NAME = normalized\n    _HUBER_DELTA = max(1e-8, float(huber_delta))\n\n\ndef regression_loss_name() -> str:\n    return str(_REGRESSION_LOSS_NAME)\n\n\ndef regression_huber_delta() -> float:\n    return float(_HUBER_DELTA)\n\n\ndef pointwise_regression_loss(pred: torch.Tensor, y: torch.Tensor) -> torch.Tensor:\n    if regression_loss_name() == \"huber\":\n        return torch.nn.functional.huber_loss(pred, y, reduction=\"none\", delta=regression_huber_delta())\n    return (pred - y) ** 2\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, model=None, grad_clip_norm: float = 0.0) -> None:\n        if grad_clip_norm > 0.0 and model is not None:\n            self.accelerator.clip_grad_norm_(model.parameters(), grad_clip_norm)\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, model=None, grad_clip_norm: float = 0.0) -> None:\n        if self.use_amp:\n            self.scaler.unscale_(optimizer)\n        if grad_clip_norm > 0.0 and model is not None:\n            torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip_norm)\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] = {}\nDEFAULT_MIN_TIMECODE = 93_000_000\n_USE_SOURCE_RAW_WEIGHTS = False\n\n\ndef configure_split_view_weight_loading(use_source_raw_weights: bool) -> None:\n    global _USE_SOURCE_RAW_WEIGHTS\n    _USE_SOURCE_RAW_WEIGHTS = bool(use_source_raw_weights)\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    if bool(_USE_SOURCE_RAW_WEIGHTS):\n        try:\n            w = np.asarray(load_split_view_source_field_slice(day_dir, \"w\"), dtype=np.float32)\n        except Exception:\n            w = np.memmap(day_dir / \"w_float32.memmap\", mode=\"r\", dtype=np.float32, shape=(n_rows,))\n    else:\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 filter_valid_end_indices(\n    day_dir: Path,\n    end_indices: np.ndarray,\n    min_timecode: int = -1,\n    require_positive_weight: bool = False,\n) -> np.ndarray:\n    ends = np.asarray(end_indices, dtype=np.int64)\n    if ends.size == 0:\n        return ends\n    mask = np.ones((int(ends.shape[0]),), dtype=bool)\n    if bool(require_positive_weight):\n        meta = load_day_meta(day_dir)\n        n_rows = int(meta[\"n_rows\"])\n        cache_w = np.memmap(day_dir / \"w_float32.memmap\", mode=\"r\", dtype=np.float32, shape=(n_rows,))\n        weight_mask = np.asarray(cache_w[ends], dtype=np.float32)\n        mask &= np.isfinite(weight_mask) & (weight_mask > 0.0)\n    if int(min_timecode) > 0:\n        dt_slice = np.asarray(load_split_view_source_field_slice(day_dir, \"datetime\"), dtype=np.int64)\n        dt_end = np.asarray(dt_slice[ends], dtype=np.int64)\n        mask &= dt_end >= int(min_timecode)\n    return ends[mask]\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(\n    cache_dir: Path,\n    day_dir: Path,\n    seq_len: int,\n    sample_stride: int,\n    min_timecode: int = -1,\n    require_positive_weight: bool = False,\n) -> 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            str(int(min_timecode)),\n            str(int(bool(require_positive_weight))),\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    min_timecode: int = -1,\n    require_positive_weight: bool = False,\n) -> np.ndarray:\n    if cache_dir is None:\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        return filter_valid_end_indices(\n            day_dir,\n            ends,\n            min_timecode=min_timecode,\n            require_positive_weight=require_positive_weight,\n        )\n    cache_path = build_end_index_cache_path(\n        cache_dir,\n        day_dir,\n        seq_len=seq_len,\n        sample_stride=sample_stride,\n        min_timecode=min_timecode,\n        require_positive_weight=require_positive_weight,\n    )\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    ends = filter_valid_end_indices(\n        day_dir,\n        ends,\n        min_timecode=min_timecode,\n        require_positive_weight=require_positive_weight,\n    )\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        min_timecode: int = -1,\n        require_positive_weight: bool = False,\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.min_timecode = int(min_timecode)\n        self.require_positive_weight = bool(require_positive_weight)\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                min_timecode=self.min_timecode,\n                require_positive_weight=self.require_positive_weight,\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        \n        # Optimize for sample_stride > 1: materialize sequences on CPU to reduce GPU transfer\n        cpu_materialize_threshold = 2.5\n        span_rows = span_end - span_start + 1\n        batch_rows = len(batch_end)\n        use_cpu_materialize = (\n            span_rows > int(batch_rows * self.seq_len * cpu_materialize_threshold)\n        )\n        \n        if use_cpu_materialize:\n            # CPU-side materialization: directly build (batch_size, seq_len, feat_dim) sequences\n            span_x_dtype = np.float32 if store.need_clean else store.x.dtype\n            batch_size = len(batch_end)\n            feat_dim = store.x.shape[1]\n            materialized_x = np.empty((batch_size, self.seq_len, feat_dim), dtype=span_x_dtype)\n            \n            for i, end_idx in enumerate(batch_end):\n                seq_start = int(end_idx) - self.seq_len + 1\n                seq_slice = store.x[seq_start : int(end_idx) + 1]\n                materialized_x[i, :, :] = seq_slice\n            \n            if store.need_clean:\n                np.nan_to_num(materialized_x, copy=False, nan=0.0, posinf=0.0, neginf=0.0)\n            \n            # For CPU-materialized case, span_x IS the materialized sequences, local_end is not needed for gather\n            span_xb = torch.from_numpy(np.ascontiguousarray(materialized_x))\n            local_endb = torch.empty((0,), dtype=torch.int64)  # sentinel: empty means already materialized\n        else:\n            # Original span-based approach for small strides\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            span_xb = torch.from_numpy(np.ascontiguousarray(span_x))\n            local_endb = torch.from_numpy(local_end)\n        \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        np.nan_to_num(y, copy=False, nan=0.0, posinf=0.0, neginf=0.0)\n        np.nan_to_num(w, copy=False, nan=0.0, posinf=0.0, neginf=0.0)\n        np.maximum(w, 0.0, out=w)\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            if local_endb.numel() > 0:\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        \n        # Handle CPU-materialized case differently\n        if part.local_end.numel() == 0:\n            # span_x is already (batch_size, seq_len, feat_dim), pad along batch dim\n            span_x = torch.cat([part.span_x, part.span_x[-1:].repeat(int(pad_size), 1, 1)], dim=0)\n            local_end = part.local_end  # keep empty sentinel\n        else:\n            # Original span-based: only pad local_end, y, w\n            span_x = part.span_x\n            local_end = torch.cat([part.local_end, part.local_end[-1:].repeat(int(pad_size))], dim=0)\n        \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        \n        if self.pin_memory:\n            span_x = span_x.pin_memory()\n            if local_end.numel() > 0:\n                local_end = local_end.pin_memory()\n            y = y.pin_memory()\n            w = w.pin_memory()\n        return PackedSeqBatchPart(span_x=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 * pointwise_regression_loss(pred, y)).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_unweighted_ic\": float(candidate_val_metrics.get(\"unweighted_ic\", candidate_val_metrics[\"ic\"])),\n        \"val_weighted_ic\": float(candidate_val_metrics.get(\"weighted_ic\", 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.get(\"val_weighted_ic\", 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        else:\n            ww = None\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        if ww is None:\n            weight = mask.to(dtype=torch.float64)\n        else:\n            weight = torch.where(mask, ww, zero).to(dtype=torch.float64)\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                weight.sum(),\n                (weight * p).sum(),\n                (weight * t).sum(),\n                (weight * p * p).sum(),\n                (weight * t * t).sum(),\n                (weight * p * t).sum(),\n                (weight * torch.abs(d)).sum(),\n                (weight * 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((16,), dtype=torch.float64, device=device)\n        return self.buf.to(device=device, dtype=torch.float64)\n\n\ndef corr_from_sums(sum_x: float, sum_y: float, sum_xx: float, sum_yy: float, sum_xy: float, denom_weight: float) -> float:\n    if (not np.isfinite(denom_weight)) or denom_weight <= 0:\n        return float(\"nan\")\n    mean_x = sum_x / denom_weight\n    mean_y = sum_y / denom_weight\n    var_x = max(sum_xx / denom_weight - mean_x * mean_x, 1e-12)\n    var_y = max(sum_yy / denom_weight - mean_y * mean_y, 1e-12)\n    cov_xy = sum_xy / denom_weight - mean_x * mean_y\n    return float(cov_xy / math.sqrt(var_x * var_y))\n\n\ndef metrics_from_tensor(t: torch.Tensor) -> dict:\n    values = [float(x) for x in t.tolist()]\n    if len(values) < 8:\n        raise ValueError(f\"metrics_from_tensor expects at least 8 values, got {len(values)}\")\n    n, sum_p, sum_y, sum_pp, sum_yy, sum_py, sum_abs, sum_sq = values[:8]\n    if (not np.isfinite(n)) or n <= 0:\n        return {\n            \"mse\": float(\"nan\"),\n            \"rmse\": float(\"nan\"),\n            \"mae\": float(\"nan\"),\n            \"ic\": float(\"nan\"),\n            \"unweighted_mse\": float(\"nan\"),\n            \"unweighted_rmse\": float(\"nan\"),\n            \"unweighted_mae\": float(\"nan\"),\n            \"unweighted_ic\": float(\"nan\"),\n            \"weighted_mse\": float(\"nan\"),\n            \"weighted_rmse\": float(\"nan\"),\n            \"weighted_mae\": float(\"nan\"),\n            \"weighted_ic\": float(\"nan\"),\n            \"weight_sum\": 0.0,\n            \"n\": 0,\n        }\n    mse = sum_sq / n\n    mae = sum_abs / n\n    ic = corr_from_sums(sum_p, sum_y, sum_pp, sum_yy, sum_py, n)\n    if len(values) >= 16:\n        weight_sum, sum_wp, sum_wy, sum_wpp, sum_wyy, sum_wpy, sum_wabs, sum_wsq = values[8:16]\n    elif len(values) >= 14:\n        weight_sum, sum_wp, sum_wy, sum_wpp, sum_wyy, sum_wpy = values[8:14]\n        sum_wabs, sum_wsq = sum_abs, sum_sq\n    else:\n        weight_sum, sum_wp, sum_wy, sum_wpp, sum_wyy, sum_wpy = n, sum_p, sum_y, sum_pp, sum_yy, sum_py\n        sum_wabs, sum_wsq = sum_abs, sum_sq\n    weighted_ic = corr_from_sums(sum_wp, sum_wy, sum_wpp, sum_wyy, sum_wpy, weight_sum)\n    weighted_mse = (sum_wsq / weight_sum) if weight_sum > 0 else float(\"nan\")\n    weighted_mae = (sum_wabs / weight_sum) if weight_sum > 0 else float(\"nan\")\n    weighted_rmse = math.sqrt(weighted_mse) if np.isfinite(weighted_mse) and weighted_mse >= 0 else float(\"nan\")\n    return {\n        \"mse\": float(weighted_mse),\n        \"rmse\": float(weighted_rmse),\n        \"mae\": float(weighted_mae),\n        \"ic\": float(weighted_ic),\n        \"unweighted_mse\": float(mse),\n        \"unweighted_rmse\": float(math.sqrt(mse)),\n        \"unweighted_mae\": float(mae),\n        \"unweighted_ic\": float(ic),\n        \"weighted_mse\": float(weighted_mse),\n        \"weighted_rmse\": float(weighted_rmse),\n        \"weighted_mae\": float(weighted_mae),\n        \"weighted_ic\": float(weighted_ic),\n        \"weight_sum\": float(weight_sum),\n        \"n\": int(n),\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    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            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    stats_t = accelerator.reduce(stats.to_tensor(accelerator.device), reduction=\"sum\")\n    m = metrics_from_tensor(stats_t)\n    m[\"loss\"] = float(m.get(\"weighted_mse\", float(\"nan\")))\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        yb = part.y.to(device, non_blocking=copy_non_blocking)\n        wb = part.w.to(device, non_blocking=copy_non_blocking)\n        \n        # Check if part is already CPU-materialized (local_end.numel() == 0 is sentinel)\n        if part.local_end.numel() == 0:\n            # Already materialized on CPU: span_x is (batch_size, seq_len, feat_dim)\n            xb = part.span_x.to(device, non_blocking=copy_non_blocking).float()\n            xb = (xb - feat_mean) / feat_std\n            xb = xb.contiguous()\n        else:\n            # Original span-based gather on GPU\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            seq_idx = local_end[:, None] + seq_offsets[None, :]\n            xb = span_x[seq_idx].float()\n            xb = (xb - feat_mean) / feat_std\n            xb = xb.contiguous()\n        \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    configure_regression_loss(getattr(args, \"regression_loss\", \"mse\"), getattr(args, \"huber_delta\", 1.0))\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    configure_split_view_weight_loading(bool(args.use_source_raw_weights))\n    min_timecode = int(args.min_timecode)\n    require_positive_weight = bool(args.require_positive_weight)\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    print(\n        f\"[sample-filter] min_timecode={min_timecode} require_positive_weight={int(require_positive_weight)} \"\n        f\"use_source_raw_weights={int(bool(args.use_source_raw_weights))}\",\n        flush=True,\n    )\n    print(\n        f\"[loss] regression_loss={regression_loss_name()} huber_delta={regression_huber_delta():.6f}\",\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        min_timecode=min_timecode,\n        require_positive_weight=require_positive_weight,\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            min_timecode=min_timecode,\n            require_positive_weight=require_positive_weight,\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            min_timecode=min_timecode,\n            require_positive_weight=require_positive_weight,\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            min_timecode=min_timecode,\n            require_positive_weight=require_positive_weight,\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    init_checkpoint = str(getattr(args, \"init_checkpoint\", \"\") or \"\").strip()\n    init_checkpoint_path: Path | None = None\n    if init_checkpoint:\n        init_checkpoint_path = Path(init_checkpoint).expanduser()\n        if not init_checkpoint_path.is_absolute():\n            init_checkpoint_path = (Path.cwd() / init_checkpoint_path).resolve()\n        if not init_checkpoint_path.is_file():\n            raise FileNotFoundError(f\"init checkpoint does not exist: {init_checkpoint_path}\")\n        state_dict = torch.load(str(init_checkpoint_path), map_location=\"cpu\", weights_only=True)\n        model.load_state_dict(state_dict)\n        if accelerator.is_main_process:\n            print(f\"[init-checkpoint] loaded={init_checkpoint_path}\", flush=True)\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_unweighted_ic\",\n                    \"train_weighted_ic\",\n                    \"train_rmse\",\n                    \"train_unweighted_rmse\",\n                    \"val_loss\",\n                    \"val_ic\",\n                    \"val_unweighted_ic\",\n                    \"val_weighted_ic\",\n                    \"val_rmse\",\n                    \"val_unweighted_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_start = max(1, int(getattr(args, \"epoch_start\", 1)))\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(epoch_start)\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                \"min_timecode\": int(min_timecode),\n                \"require_positive_weight\": bool(require_positive_weight),\n                \"use_source_raw_weights\": bool(args.use_source_raw_weights),\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 local_epoch in range(1, epochs + 1):\n        epoch = epoch_start + local_epoch - 1\n        current_lr = float(epoch_lrs[local_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 local_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, model=model, grad_clip_norm=grad_clip_norm)\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                \"unweighted_ic\": float(\"nan\"),\n                \"weighted_ic\": float(\"nan\"),\n                \"rmse\": float(\"nan\"),\n                \"mae\": float(\"nan\"),\n                \"weight_sum\": 0.0,\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        current_weighted_ic = float(val_m.get(\"weighted_ic\", val_m[\"ic\"]))\n        best_weighted_ic = float(best_val_m.get(\"weighted_ic\", best_val_m[\"ic\"])) if best_val_m is not None else float(\"-inf\")\n        if best_val_m is None or current_weighted_ic > best_weighted_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.get(\"unweighted_ic\", float(\"nan\")),\n                        train_m.get(\"weighted_ic\", float(\"nan\")),\n                        train_m[\"rmse\"],\n                        train_m.get(\"unweighted_rmse\", float(\"nan\")),\n                        val_m[\"loss\"],\n                        val_m[\"ic\"],\n                        val_m.get(\"unweighted_ic\", float(\"nan\")),\n                        val_m.get(\"weighted_ic\", float(\"nan\")),\n                        val_m[\"rmse\"],\n                        val_m.get(\"unweighted_rmse\", float(\"nan\")),\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} \"\n                f\"train_uic={train_m.get('unweighted_ic', float('nan')):.6f} train_rmse={train_m['rmse']:.6f} \"\n                f\"val_loss={val_m['loss']:.6f} val_ic={val_m['ic']:.6f} \"\n                f\"val_uic={val_m.get('unweighted_ic', float('nan')):.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                f\"val_wic={checkpoint_entry.get('val_weighted_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} \"\n                f\"ic={test_m_best['ic']:.6f} uic={test_m_best.get('unweighted_ic', float('nan')):.6f} \"\n                f\"wic={test_m_best.get('weighted_ic', float('nan')):.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} \"\n                    f\"ic={test_m_selected['ic']:.6f} uic={test_m_selected.get('unweighted_ic', float('nan')):.6f} \"\n                    f\"wic={test_m_selected.get('weighted_ic', float('nan')):.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            \"epoch_start\": int(epoch_start),\n            \"epochs_this_run\": int(epochs),\n            \"init_checkpoint\": str(init_checkpoint_path) if init_checkpoint_path is not None else \"\",\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_weighted_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            \"min_timecode\": int(min_timecode),\n            \"require_positive_weight\": bool(require_positive_weight),\n            \"use_source_raw_weights\": bool(args.use_source_raw_weights),\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\": f\"weighted_{regression_loss_name()}\",\n            \"loss_config\": {\n                \"regression_loss\": regression_loss_name(),\n                \"huber_delta\": regression_huber_delta(),\n            },\n            \"metric_main\": \"val_weighted_ic\",\n            \"test_selection_rule\": f\"val_weighted_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_epoch_by_val_weighted_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                \"min_timecode\": int(min_timecode),\n                \"require_positive_weight\": bool(require_positive_weight),\n                \"use_source_raw_weights\": bool(args.use_source_raw_weights),\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(\"--regression-loss\", type=str, default=\"mse\", choices=[\"mse\", \"huber\"])\n    p.add_argument(\"--huber-delta\", type=float, default=1.0)\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(\n        \"--use-source-raw-weights\",\n        type=int,\n        default=1,\n        help=\"1 to load raw continuous weights from the source split npy instead of cache-side binary masks.\",\n    )\n    p.add_argument(\n        \"--min-timecode\",\n        type=int,\n        default=DEFAULT_MIN_TIMECODE,\n        help=\"Only keep sequence end points whose source datetime >= this HHMMSSmmm timecode.\",\n    )\n    p.add_argument(\n        \"--require-positive-weight\",\n        type=int,\n        default=1,\n        help=\"1 to drop sequence end points whose cached/source weight is not positive.\",\n    )\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    p.add_argument(\"--init-checkpoint\", type=str, default=\"\")\n    p.add_argument(\"--epoch-start\", 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}}}}