"""Run forecasting-v4 masked continuous-horizon experiments.""" from __future__ import annotations import argparse import hashlib import json import math import platform import random import sys import time from datetime import datetime, timezone from pathlib import Path from typing import Any, Iterable import numpy as np import torch from torch import nn from project.data.close_distribution_v2 import ( BASE_TOKEN_FEATURE_NAMES, CANONICAL_HORIZONS, CLASS_COUNT, CONFIRMATION_ASSETS, DISCOVERY_FOLDS, EVALUATION_HORIZONS, SELECTION_ASSETS, TOKEN_FEATURE_NAMES, ) from project.evaluators.close_distribution_v2 import ( ClassMap, class_map_to_arrays, encode_returns, fit_class_map, metrics, prior_probabilities, score_arrays, ) from project.models.close_distribution_v2 import build_model MODEL_NAMES = ( "prior", "pooled_linear", "tcn", "transformer_learned", "transformer_scalar", "transformer_rotary", ) CONTINUOUS_MODELS = ("transformer_scalar", "transformer_rotary") OPTIMIZER_RECIPES = ( "adamw_constant", "adamw_cosine", "adamw_warmup_cosine", "lion_warmup_cosine", ) SESSION_TIME_FEATURE_NAMES = ( "regular_session_progress", "regular_session_progress_sin", "regular_session_progress_cos", ) SESSION_GAP_FEATURE_NAMES = ( "log1p_minutes_since_previous_session_last_observed_close", ) RELATIVE_PRICE_FEATURE_NAMES = ( "close_log_return_since_session_first_observed", "close_log_return_since_previous_session_last_observed", ) RAW_PRICE_FEATURE_NAMES = ("log_raw_close",) FEATURE_SETS = { "baseline": BASE_TOKEN_FEATURE_NAMES, "session_time": BASE_TOKEN_FEATURE_NAMES + SESSION_TIME_FEATURE_NAMES, "session_gap": BASE_TOKEN_FEATURE_NAMES + SESSION_GAP_FEATURE_NAMES, "relative_price": ( BASE_TOKEN_FEATURE_NAMES + RELATIVE_PRICE_FEATURE_NAMES ), "raw_price": BASE_TOKEN_FEATURE_NAMES + RAW_PRICE_FEATURE_NAMES, "session_time_gap": ( BASE_TOKEN_FEATURE_NAMES + SESSION_TIME_FEATURE_NAMES + SESSION_GAP_FEATURE_NAMES ), "session_time_relative": ( BASE_TOKEN_FEATURE_NAMES + SESSION_TIME_FEATURE_NAMES + RELATIVE_PRICE_FEATURE_NAMES ), "session_gap_relative": ( BASE_TOKEN_FEATURE_NAMES + SESSION_GAP_FEATURE_NAMES + RELATIVE_PRICE_FEATURE_NAMES ), "all_without_raw": tuple( name for name in TOKEN_FEATURE_NAMES if name != "log_raw_close" ), "all_natural": TOKEN_FEATURE_NAMES, "all_without_session_time": tuple( name for name in TOKEN_FEATURE_NAMES if name not in SESSION_TIME_FEATURE_NAMES ), "all_without_session_gap": tuple( name for name in TOKEN_FEATURE_NAMES if name not in SESSION_GAP_FEATURE_NAMES ), "all_without_relative_price": tuple( name for name in TOKEN_FEATURE_NAMES if name not in RELATIVE_PRICE_FEATURE_NAMES ), } SPARSE_FEATURES = { "close_log_return_since_last_observed_close", "log_raw_close", "close_log_return_since_session_first_observed", "close_log_return_since_previous_session_last_observed", } PROVENANCE_PATHS = ( "project/data/close_distribution_v2.py", "project/evaluators/close_distribution_v2.py", "project/experiments/forecasting_v4/runner.py", "project/models/close_distribution_v2.py", "project/studies/forecasting_v4.md", "requirements.txt", ) class Lion(torch.optim.Optimizer): """Minimal Lion optimizer used for the frozen optimizer treatment.""" def __init__( self, params: Iterable[torch.Tensor], *, lr: float, betas: tuple[float, float] = (0.9, 0.99), weight_decay: float = 0.0, ) -> None: if lr <= 0.0: raise ValueError("Lion learning rate must be positive") if not 0.0 <= weight_decay: raise ValueError("Lion weight decay cannot be negative") if not all(0.0 <= beta < 1.0 for beta in betas): raise ValueError("Lion betas must lie in [0, 1)") super().__init__( params, { "lr": lr, "betas": betas, "weight_decay": weight_decay, }, ) @torch.no_grad() def step(self, closure: Any = None) -> Any: loss = None if closure is None else closure() for group in self.param_groups: beta1, beta2 = group["betas"] for parameter in group["params"]: if parameter.grad is None: continue gradient = parameter.grad if gradient.is_sparse: raise RuntimeError("Lion does not support sparse gradients") if group["weight_decay"]: parameter.mul_(1.0 - group["lr"] * group["weight_decay"]) state = self.state[parameter] if not state: state["exp_avg"] = torch.zeros_like(parameter) average = state["exp_avg"] update = average.mul(beta1).add( gradient, alpha=1.0 - beta1, ) parameter.add_(torch.sign(update), alpha=-group["lr"]) average.mul_(beta2).add_(gradient, alpha=1.0 - beta2) return loss def _ordinal(value: str) -> np.int32: return np.datetime64(value, "D").astype(np.int32) def _sha256(path: Path) -> str: digest = hashlib.sha256() with path.open("rb") as stream: for chunk in iter(lambda: stream.read(8 * 1024 * 1024), b""): digest.update(chunk) return digest.hexdigest() def _load_npz(path: Path) -> dict[str, np.ndarray]: with np.load(path, allow_pickle=False) as bundle: return {name: bundle[name] for name in bundle.files} def _configure(seed: int, threads: int) -> None: random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.set_num_threads(threads) torch.use_deterministic_algorithms(True) def _parse_horizons(value: str | Iterable[int]) -> tuple[int, ...]: if isinstance(value, str): result = tuple(int(item.strip()) for item in value.split(",")) else: result = tuple(int(item) for item in value) allowed = set(int(item) for item in EVALUATION_HORIZONS) if ( not result or len(set(result)) != len(result) or any(item not in allowed for item in result) or tuple(sorted(result)) != result ): raise ValueError("supervised horizons must be unique sorted values 2..32") return result def _asset_manifest( study_dir: Path, asset: str, ) -> tuple[dict[str, Any], Path]: if asset in SELECTION_ASSETS: manifest_path = study_dir / "data" / "prepared-manifest.json" data_dir = study_dir / "data" / "runner" elif asset in CONFIRMATION_ASSETS: marker_path = study_dir / "confirmation" / "CONFIRMATION_OPENED.json" if not marker_path.exists(): raise PermissionError("confirmation asset is still protected") marker = json.loads(marker_path.read_text()) manifest_path = ( study_dir / "data" / "protected" / "confirmation" / "prepared-manifest.json" ) data_dir = manifest_path.parent manifest = json.loads(manifest_path.read_text()) if manifest.get("freeze_sha256") != marker.get("freeze_sha256"): raise ValueError("confirmation manifest does not match the freeze") return manifest, data_dir else: raise ValueError("asset is outside the frozen v4 whitelist") manifest = json.loads(manifest_path.read_text()) if manifest.get("scope") != "selection": raise ValueError("selection manifest has an invalid scope") return manifest, data_dir def _feature_indices(feature_set: str) -> tuple[int, ...]: names = FEATURE_SETS[feature_set] return tuple(TOKEN_FEATURE_NAMES.index(name) for name in names) def _normalization( features: dict[str, np.ndarray], *, train_end: str, selected_indices: tuple[int, ...], max_row_exclusive: int | None = None, ) -> tuple[np.ndarray, np.ndarray]: train = features["session_date"] < _ordinal(train_end) if max_row_exclusive is not None: train &= np.arange(len(train)) < max_row_exclusive observed = features["X"][:, 1] > 0.5 center = np.empty(len(selected_indices), dtype=np.float32) scale = np.empty(len(selected_indices), dtype=np.float32) for output_index, feature_index in enumerate(selected_indices): name = TOKEN_FEATURE_NAMES[feature_index] if name == "close_return_observed": center[output_index] = 0.0 scale[output_index] = 1.0 continue selected = train & observed if name in SPARSE_FEATURES else train values = features["X"][selected, feature_index].astype(np.float64) if not len(values) or not np.all(np.isfinite(values)): raise ValueError(f"invalid normalization values for {name}") center[output_index] = np.mean(values) scale[output_index] = max(float(np.std(values)), 1e-6) return center, scale def _context( x: np.ndarray, rows: np.ndarray, *, context_length: int, selected_indices: tuple[int, ...], center: np.ndarray, scale: np.ndarray, ) -> tuple[torch.Tensor, torch.Tensor]: offsets = np.arange(1 - context_length, 1, dtype=np.int64) indices = rows[:, None] + offsets[None, :] valid = indices >= 0 safe = np.maximum(indices, 0) values = np.asarray( x[safe][..., selected_indices], dtype=np.float32, ).copy() values = (values - center[None, None, :]) / scale[None, None, :] observed = x[safe, 1] > 0.5 for local_index, feature_index in enumerate(selected_indices): if TOKEN_FEATURE_NAMES[feature_index] in SPARSE_FEATURES: values[..., local_index] *= observed values *= valid[:, :, None] values = np.concatenate( (values, valid[:, :, None].astype(np.float32)), axis=2, ) return torch.from_numpy(values), torch.from_numpy(~valid) def _optimizer_settings(recipe: str) -> tuple[float, float]: if recipe in ("adamw_constant", "adamw_cosine"): return 0.002, 0.0001 if recipe == "adamw_warmup_cosine": return 0.004, 0.0001 if recipe == "lion_warmup_cosine": return 0.0004, 0.0005 raise ValueError(f"unknown optimizer recipe: {recipe}") def _optimizer( model: nn.Module, recipe: str, ) -> torch.optim.Optimizer: learning_rate, weight_decay = _optimizer_settings(recipe) if recipe.startswith("adamw"): return torch.optim.AdamW( model.parameters(), lr=learning_rate, weight_decay=weight_decay, ) return Lion( model.parameters(), lr=learning_rate, weight_decay=weight_decay, ) def _scheduled_learning_rate( recipe: str, *, update: int, total_updates: int, ) -> float: peak, _weight_decay = _optimizer_settings(recipe) if recipe == "adamw_constant": return peak warmup_updates = ( max(1, math.ceil(0.05 * total_updates)) if "warmup" in recipe else 0 ) if update < warmup_updates: return peak * (update + 1) / warmup_updates remaining = max(total_updates - warmup_updates, 1) progress = min( max((update - warmup_updates) / max(remaining - 1, 1), 0.0), 1.0, ) return peak * 0.5 * (1.0 + math.cos(math.pi * progress)) def _masked_loss_part( logits: torch.Tensor, labels: torch.Tensor, mask: torch.Tensor, denominators: torch.Tensor, ) -> torch.Tensor: losses = nn.functional.cross_entropy( logits.reshape(-1, logits.shape[-1]), labels.clamp_min(0).reshape(-1), reduction="none", ).reshape(labels.shape) horizon_sum = (losses * mask).sum(dim=0) return torch.mean(horizon_sum / denominators) def _masked_nll( probabilities: np.ndarray, labels: np.ndarray, target_mask: np.ndarray, ) -> float: p = np.asarray(probabilities, dtype=np.float64) y = np.asarray(labels, dtype=np.int64) mask = np.asarray(target_mask, dtype=np.bool_) if p.shape[:2] != y.shape or y.shape != mask.shape: raise ValueError("diagnostic arrays do not align") by_horizon = [] for horizon in range(y.shape[1]): selected = mask[:, horizon] if not np.any(selected): raise ValueError("diagnostic horizon has no valid target") selected_probability = p[ selected, horizon, y[selected, horizon], ] by_horizon.append( -np.mean(np.log(np.clip(selected_probability, 1e-12, 1.0))) ) return float(np.mean(by_horizon)) def _train_model( model: nn.Module, *, x: np.ndarray, row_index: np.ndarray, labels: np.ndarray, target_mask: np.ndarray, context_length: int, selected_indices: tuple[int, ...], center: np.ndarray, scale: np.ndarray, seed: int, epochs: int, batch_size: int, gradient_accumulation_steps: int, optimizer_recipe: str, device: torch.device, diagnostic_rows: np.ndarray, diagnostic_labels: np.ndarray, diagnostic_target_mask: np.ndarray, inference_batch_size: int, early_stopping_patience: int, early_stopping_min_delta: float, ) -> tuple[ list[float], list[float], int, list[float], int, bool, ]: if early_stopping_patience < 0: raise ValueError("early stopping patience cannot be negative") if early_stopping_min_delta < 0.0: raise ValueError("early stopping minimum delta cannot be negative") if early_stopping_patience and not len(diagnostic_rows): raise ValueError("early stopping requires a diagnostic tail") optimizer = _optimizer(model, optimizer_recipe) effective_batch = batch_size * gradient_accumulation_steps updates_per_epoch = math.ceil(len(row_index) / effective_batch) total_updates = updates_per_epoch * epochs generator = torch.Generator().manual_seed(seed + 10_007) history: list[float] = [] learning_rates: list[float] = [] diagnostic_history: list[float] = [] update_count = 0 best_epoch = 0 best_diagnostic_nll = float("inf") plateau_reference_nll = float("inf") best_state: dict[str, torch.Tensor] | None = None stale_epochs = 0 stopped_early = False for _epoch in range(epochs): model.train() order = torch.randperm(len(row_index), generator=generator).numpy() epoch_loss_sum = np.zeros(labels.shape[1], dtype=np.float64) epoch_target_count = np.zeros(labels.shape[1], dtype=np.int64) for start in range(0, len(order), effective_batch): update_order = order[start : start + effective_batch] update_mask = target_mask[update_order] denominators = torch.from_numpy( update_mask.sum(axis=0).astype(np.float32) ).to(device) if torch.any(denominators == 0): raise ValueError( "an optimizer update has no target for a horizon" ) optimizer.zero_grad(set_to_none=True) for micro_start in range(0, len(update_order), batch_size): batch = update_order[micro_start : micro_start + batch_size] tokens, padding = _context( x, row_index[batch], context_length=context_length, selected_indices=selected_indices, center=center, scale=scale, ) batch_labels = torch.from_numpy( labels[batch].astype(np.int64, copy=False) ).to(device) batch_mask = torch.from_numpy( target_mask[batch].astype(np.float32, copy=False) ).to(device) logits = model(tokens.to(device), padding.to(device)) loss = _masked_loss_part( logits, batch_labels, batch_mask, denominators, ) loss.backward() with torch.no_grad(): item_losses = nn.functional.cross_entropy( logits.reshape(-1, logits.shape[-1]), batch_labels.clamp_min(0).reshape(-1), reduction="none", ).reshape(batch_labels.shape) epoch_loss_sum += ( item_losses * batch_mask ).sum(dim=0).cpu().numpy() epoch_target_count += batch_mask.sum( dim=0 ).cpu().numpy().astype(np.int64) learning_rate = _scheduled_learning_rate( optimizer_recipe, update=update_count, total_updates=total_updates, ) for group in optimizer.param_groups: group["lr"] = learning_rate nn.utils.clip_grad_norm_(model.parameters(), 1.0) optimizer.step() learning_rates.append(learning_rate) update_count += 1 history.append( float(np.mean(epoch_loss_sum / epoch_target_count)) ) if len(diagnostic_rows): diagnostic_probabilities = _predict( model, x=x, row_index=diagnostic_rows, context_length=context_length, selected_indices=selected_indices, center=center, scale=scale, batch_size=inference_batch_size, device=device, query_horizons=None, ) diagnostic_nll = _masked_nll( diagnostic_probabilities, diagnostic_labels, diagnostic_target_mask, ) diagnostic_history.append(diagnostic_nll) if diagnostic_nll < best_diagnostic_nll: best_diagnostic_nll = diagnostic_nll best_epoch = _epoch + 1 best_state = { name: value.detach().cpu().clone() for name, value in model.state_dict().items() } if ( diagnostic_nll < plateau_reference_nll - early_stopping_min_delta ): plateau_reference_nll = diagnostic_nll stale_epochs = 0 else: stale_epochs += 1 if ( early_stopping_patience and stale_epochs >= early_stopping_patience ): stopped_early = True break if best_state is not None: model.load_state_dict(best_state) elif not len(diagnostic_rows): best_epoch = len(history) return ( history, learning_rates, update_count, diagnostic_history, best_epoch, stopped_early, ) def _predict( model: nn.Module, *, x: np.ndarray, row_index: np.ndarray, context_length: int, selected_indices: tuple[int, ...], center: np.ndarray, scale: np.ndarray, batch_size: int, device: torch.device, query_horizons: tuple[int, ...] | None, ) -> np.ndarray: horizon_count = ( len(query_horizons) if query_horizons is not None else len(model.output_horizons) if hasattr(model, "output_horizons") else model.horizon_count ) result = np.empty( (len(row_index), horizon_count, CLASS_COUNT), dtype=np.float64, ) query = ( torch.tensor(query_horizons, dtype=torch.float32, device=device) if query_horizons is not None else None ) model.eval() with torch.inference_mode(): for start in range(0, len(row_index), batch_size): stop = min(start + batch_size, len(row_index)) tokens, padding = _context( x, row_index[start:stop], context_length=context_length, selected_indices=selected_indices, center=center, scale=scale, ) logits = model( tokens.to(device), padding.to(device), query, ) result[start:stop] = torch.softmax( logits.cpu().to(torch.float64), dim=2, ).numpy() return result def interpolate_log_probabilities( supervised_probabilities: np.ndarray, supervised_horizons: Iterable[int], query_horizons: Iterable[int], ) -> np.ndarray: probabilities = np.asarray(supervised_probabilities, dtype=np.float64) supervised = np.asarray(tuple(supervised_horizons), dtype=np.float64) query = np.asarray(tuple(query_horizons), dtype=np.float64) if ( probabilities.ndim != 3 or probabilities.shape[1] != len(supervised) or np.any(np.diff(supervised) <= 0.0) or query.min() < supervised.min() or query.max() > supervised.max() ): raise ValueError("log-probability interpolation inputs are invalid") coordinate = np.log2(supervised / 2.0) query_coordinate = np.log2(query / 2.0) log_probability = np.log(np.clip(probabilities, 1e-12, 1.0)) result = np.empty( (len(probabilities), len(query), probabilities.shape[2]), dtype=np.float64, ) for output_index, value in enumerate(query_coordinate): exact = np.flatnonzero(np.isclose(coordinate, value)) if len(exact): result[:, output_index] = probabilities[:, exact[0]] continue upper = int(np.searchsorted(coordinate, value, side="right")) lower = upper - 1 weight = (value - coordinate[lower]) / ( coordinate[upper] - coordinate[lower] ) interpolated = ( (1.0 - weight) * log_probability[:, lower] + weight * log_probability[:, upper] ) interpolated -= np.max(interpolated, axis=1, keepdims=True) result[:, output_index] = np.exp(interpolated) result[:, output_index] /= result[ :, output_index ].sum(axis=1, keepdims=True) return result def _horizon_indices(horizons: Iterable[int]) -> tuple[int, ...]: positions = { int(value): index for index, value in enumerate(EVALUATION_HORIZONS) } return tuple(positions[int(value)] for value in horizons) def _session_scores( returns: np.ndarray, target_mask: np.ndarray, probabilities: np.ndarray, session_date: np.ndarray, *, class_map: ClassMap, score_values: dict[str, np.ndarray] | None = None, ) -> dict[str, np.ndarray]: scores = ( score_arrays( returns, target_mask, probabilities, class_map=class_map, ) if score_values is None else score_values ) sessions = np.unique(session_date) horizon_count = np.zeros( (len(sessions), returns.shape[1]), dtype=np.int64, ) horizon_nll_sum = np.zeros_like(horizon_count, dtype=np.float64) horizon_rps_sum = np.zeros_like(horizon_nll_sum) for index, session in enumerate(sessions): selected = session_date == session horizon_count[index] = target_mask[selected].sum(axis=0) horizon_nll_sum[index] = np.nansum( scores["nll"][selected], axis=0, ) horizon_rps_sum[index] = np.nansum( scores["ranked_probability_score"][selected], axis=0, ) return { "session_date": sessions, "horizon_count": horizon_count, "horizon_nll_sum": horizon_nll_sum, "horizon_ranked_probability_score_sum": horizon_rps_sum, } def _run_manifest(config: dict[str, Any]) -> dict[str, Any]: project_root = Path(__file__).resolve().parents[3] return { "code_sha256": { relative: _sha256(project_root / relative) for relative in PROVENANCE_PATHS }, "command": sys.argv, "config": config, "created_at_utc": datetime.now(timezone.utc).isoformat(), "environment": { "machine": platform.machine(), "numpy": np.__version__, "platform": platform.platform(), "python": sys.version, "torch": torch.__version__, "torch_mps_available": torch.backends.mps.is_available(), }, } def run( *, study_dir: Path, asset: str, model_name: str, supervised_horizons: str | Iterable[int], context_length: int, seed: int, output_dir: Path, epochs: int, batch_size: int, inference_batch_size: int, target_parameters: int, torch_threads: int, max_train_rows: int, device_name: str, gradient_accumulation_steps: int, optimizer_recipe: str, rotary_base: float, feature_set: str, diagnostic_tail_fraction: float, early_stopping_patience: int, early_stopping_min_delta: float, ) -> dict[str, Any]: started = time.monotonic() horizons = _parse_horizons(supervised_horizons) if model_name not in MODEL_NAMES: raise ValueError(f"model_name must be one of {MODEL_NAMES}") if optimizer_recipe not in OPTIMIZER_RECIPES: raise ValueError( f"optimizer_recipe must be one of {OPTIMIZER_RECIPES}" ) if feature_set not in FEATURE_SETS: raise ValueError(f"feature_set must be one of {tuple(FEATURE_SETS)}") if model_name == "pooled_linear" and feature_set != "baseline": raise ValueError("pooled linear is defined only for baseline inputs") if context_length <= 0 or batch_size <= 0: raise ValueError("context and batch sizes must be positive") if epochs <= 0 or gradient_accumulation_steps <= 0: raise ValueError("epochs and accumulation must be positive") if not 0.0 <= diagnostic_tail_fraction < 0.5: raise ValueError("diagnostic tail fraction must lie in [0, 0.5)") if early_stopping_patience and diagnostic_tail_fraction == 0.0: raise ValueError("early stopping requires a diagnostic tail") if horizons[0] != 2 or horizons[-1] != 32: raise ValueError( "supervised horizons must include endpoints 2 and 32" ) if device_name == "mps" and not torch.backends.mps.is_available(): raise ValueError("MPS was requested but is unavailable") manifest, data_dir = _asset_manifest(study_dir, asset) if asset not in manifest["assets"]: raise ValueError("asset is absent from its prepared manifest") selected_indices = _feature_indices(feature_set) learning_rate, weight_decay = _optimizer_settings(optimizer_recipe) config = { "asset": asset, "batch_size": batch_size, "context_length": context_length, "device_name": device_name, "diagnostic_tail_fraction": diagnostic_tail_fraction, "early_stopping_min_delta": early_stopping_min_delta, "early_stopping_patience": early_stopping_patience, "epochs": epochs, "feature_names": list(FEATURE_SETS[feature_set]), "feature_set": feature_set, "gradient_accumulation_steps": gradient_accumulation_steps, "inference_batch_size": inference_batch_size, "learning_rate": learning_rate, "max_train_rows": max_train_rows, "model_name": model_name, "optimizer_recipe": optimizer_recipe, "rotary_base": rotary_base, "seed": seed, "snapshot_sha256": manifest["snapshot_sha256"], "supervised_horizons": list(horizons), "target_parameters": target_parameters, "torch_threads": torch_threads, "weight_decay": weight_decay, } output_dir.mkdir(parents=True, exist_ok=True) results_path = output_dir / "results.json" run_manifest_path = output_dir / "run_manifest.json" if results_path.exists() or run_manifest_path.exists(): raise FileExistsError("experiment output already exists") run_manifest_path.write_text( json.dumps(_run_manifest(config), indent=2, sort_keys=True) + "\n" ) run_manifest_sha256 = _sha256(run_manifest_path) slug = asset.lower() feature_path = data_dir / f"{slug}_features.npz" label_path = data_dir / f"{slug}_labels.npz" for path in (feature_path, label_path): relative = str(path.relative_to(study_dir)) if _sha256(path) != manifest["derived_sha256"][relative]: raise ValueError(f"prepared hash mismatch: {relative}") features = _load_npz(feature_path) bundle = _load_npz(label_path) _configure(seed, torch_threads) device = torch.device(device_name) supervised_indices = _horizon_indices(horizons) canonical_indices = _horizon_indices(CANONICAL_HORIZONS) dense_unseen_horizons = tuple( int(value) for value in EVALUATION_HORIZONS if value not in set(CANONICAL_HORIZONS.tolist()) ) dense_unseen_indices = _horizon_indices(dense_unseen_horizons) withheld_horizons = tuple( value for value in (4, 16) if value not in horizons ) withheld_indices = _horizon_indices(withheld_horizons) fold_results: dict[str, Any] = {} for fold_name, split in DISCOVERY_FOLDS.items(): fold_started = time.monotonic() dates = bundle["session_date"] train = ( (dates >= _ordinal(split["train"][0])) & (dates < _ordinal(split["train"][1])) ) validation = ( (dates >= _ordinal(split["validation"][0])) & (dates < _ordinal(split["validation"][1])) ) train_rows = bundle["row_index"][train] train_y = bundle["y"][train] train_target_mask = bundle["target_mask"][train] if max_train_rows > 0 and len(train_rows) > max_train_rows: selected = np.linspace( 0, len(train_rows) - 1, max_train_rows, dtype=np.int64, ) train_rows = train_rows[selected] train_y = train_y[selected] train_target_mask = train_target_mask[selected] diagnostic_rows = np.empty(0, dtype=train_rows.dtype) diagnostic_y = np.empty( (0, train_y.shape[1]), dtype=train_y.dtype, ) diagnostic_target_mask = np.empty( (0, train_target_mask.shape[1]), dtype=np.bool_, ) diagnostic_purge_rows = 0 normalization_max_row = None if diagnostic_tail_fraction: diagnostic_count = max( 1, round(len(train_rows) * diagnostic_tail_fraction), ) diagnostic_rows = train_rows[-diagnostic_count:] diagnostic_y = train_y[-diagnostic_count:] diagnostic_target_mask = train_target_mask[-diagnostic_count:] train_rows = train_rows[:-diagnostic_count] train_y = train_y[:-diagnostic_count] train_target_mask = train_target_mask[:-diagnostic_count] normalization_max_row = int(diagnostic_rows[0]) unpurged_count = len(train_rows) keep = train_rows + int(EVALUATION_HORIZONS.max()) < int( diagnostic_rows[0] ) train_rows = train_rows[keep] train_y = train_y[keep] train_target_mask = train_target_mask[keep] diagnostic_purge_rows = unpurged_count - len(train_rows) class_map = fit_class_map( train_y, train_target_mask, class_count=CLASS_COUNT, ) validation_rows = bundle["row_index"][validation] validation_y = bundle["y"][validation] validation_target_mask = bundle["target_mask"][validation] center, scale = _normalization( features, train_end=split["train"][1], selected_indices=selected_indices, max_row_exclusive=normalization_max_row, ) history: list[float] = [] learning_rates: list[float] = [] diagnostic_history: list[float] = [] best_epoch = 0 stopped_early = False updates = 0 parameter_count = 0 estimated_madds = 0 model = None if model_name == "prior": direct_probabilities = prior_probabilities( len(validation_y), class_map, ) supervised_probabilities = direct_probabilities[ :, supervised_indices ] else: model = build_model( model_name, context_length=context_length, channels=len(selected_indices) + 1, output_horizons=horizons, classes=CLASS_COUNT, target_parameters=target_parameters, rotary_base=rotary_base, ).to(device) parameter_count = sum( parameter.numel() for parameter in model.parameters() if parameter.requires_grad ) estimated_madds = int(model.estimated_madds) all_labels = encode_returns( train_y, train_target_mask, class_map.edges, ) diagnostic_labels = encode_returns( diagnostic_y, diagnostic_target_mask, class_map.edges, ) ( history, learning_rates, updates, diagnostic_history, best_epoch, stopped_early, ) = _train_model( model, x=features["X"], row_index=train_rows, labels=all_labels[:, supervised_indices], target_mask=train_target_mask[:, supervised_indices], context_length=context_length, selected_indices=selected_indices, center=center, scale=scale, seed=seed, epochs=epochs, batch_size=batch_size, gradient_accumulation_steps=( gradient_accumulation_steps ), optimizer_recipe=optimizer_recipe, device=device, diagnostic_rows=diagnostic_rows, diagnostic_labels=diagnostic_labels[:, supervised_indices], diagnostic_target_mask=diagnostic_target_mask[ :, supervised_indices ], inference_batch_size=inference_batch_size, early_stopping_patience=early_stopping_patience, early_stopping_min_delta=early_stopping_min_delta, ) supervised_probabilities = _predict( model, x=features["X"], row_index=validation_rows, context_length=context_length, selected_indices=selected_indices, center=center, scale=scale, batch_size=inference_batch_size, device=device, query_horizons=( horizons if model_name in CONTINUOUS_MODELS else None ), ) interpolation_probabilities = interpolate_log_probabilities( supervised_probabilities, horizons, EVALUATION_HORIZONS, ) if model_name in CONTINUOUS_MODELS: if model is None: raise AssertionError("continuous model was not constructed") direct_probabilities = _predict( model, x=features["X"], row_index=validation_rows, context_length=context_length, selected_indices=selected_indices, center=center, scale=scale, batch_size=inference_batch_size, device=device, query_horizons=tuple( int(value) for value in EVALUATION_HORIZONS ), ) elif model_name != "prior": direct_probabilities = interpolation_probabilities prior = prior_probabilities(len(validation_y), class_map) direct_scores = score_arrays( validation_y, validation_target_mask, direct_probabilities, class_map=class_map, ) interpolation_scores = score_arrays( validation_y, validation_target_mask, interpolation_probabilities, class_map=class_map, ) prior_scores = score_arrays( validation_y, validation_target_mask, prior, class_map=class_map, ) direct_metrics = metrics( validation_y, validation_target_mask, direct_probabilities, class_map=class_map, score_values=direct_scores, ) metric_slices = { "canonical": metrics( validation_y, validation_target_mask, direct_probabilities, class_map=class_map, horizon_indices=canonical_indices, score_values=direct_scores, ), "dense_unseen": metrics( validation_y, validation_target_mask, direct_probabilities, class_map=class_map, horizon_indices=dense_unseen_indices, score_values=direct_scores, ), "supervised": metrics( validation_y, validation_target_mask, direct_probabilities, class_map=class_map, horizon_indices=supervised_indices, score_values=direct_scores, ), "interpolation_baseline_dense_unseen": metrics( validation_y, validation_target_mask, interpolation_probabilities, class_map=class_map, horizon_indices=dense_unseen_indices, score_values=interpolation_scores, ), "prior_canonical": metrics( validation_y, validation_target_mask, prior, class_map=class_map, horizon_indices=canonical_indices, score_values=prior_scores, ), "prior_dense_unseen": metrics( validation_y, validation_target_mask, prior, class_map=class_map, horizon_indices=dense_unseen_indices, score_values=prior_scores, ), } if withheld_indices: metric_slices["withheld"] = metrics( validation_y, validation_target_mask, direct_probabilities, class_map=class_map, horizon_indices=withheld_indices, score_values=direct_scores, ) metric_slices["interpolation_baseline_withheld"] = metrics( validation_y, validation_target_mask, interpolation_probabilities, class_map=class_map, horizon_indices=withheld_indices, score_values=interpolation_scores, ) metric_slices["prior_withheld"] = metrics( validation_y, validation_target_mask, prior, class_map=class_map, horizon_indices=withheld_indices, score_values=prior_scores, ) np.savez_compressed( output_dir / f"{fold_name}_direct_session_scores.npz", **_session_scores( validation_y, validation_target_mask, direct_probabilities, bundle["session_date"][validation], class_map=class_map, score_values=direct_scores, ), ) np.savez_compressed( output_dir / f"{fold_name}_interpolation_session_scores.npz", **_session_scores( validation_y, validation_target_mask, interpolation_probabilities, bundle["session_date"][validation], class_map=class_map, score_values=interpolation_scores, ), ) if model is not None: torch.save( { "asset": asset, "center": center, "context_length": context_length, "feature_names": FEATURE_SETS[feature_set], "model_name": model_name, "rotary_base": rotary_base, "scale": scale, "state_dict": model.state_dict(), "supervised_horizons": horizons, **class_map_to_arrays(class_map), }, output_dir / f"{fold_name}_model.pt", ) fold_results[fold_name] = { "best_epoch": best_epoch, "class_map_rows": int(len(train_rows)), "checkpoint_selection": ( "lowest_inner_tail_nll" if len(diagnostic_rows) else "final_epoch" ), "completed_fixed_update_budget": bool( model_name == "prior" or updates == math.ceil( len(train_rows) / (batch_size * gradient_accumulation_steps) ) * epochs ), "converged_by_inner_tail": bool( model_name == "prior" or not diagnostic_tail_fraction or stopped_early ), "diagnostic_purge_rows": diagnostic_purge_rows, "diagnostic_tail_nll": diagnostic_history, "diagnostic_tail_rows": int(len(diagnostic_rows)), "direct_is_continuous_query": model_name in CONTINUOUS_MODELS, "duration_seconds": time.monotonic() - fold_started, "estimated_madds_per_example": estimated_madds, "learning_rate_first": ( learning_rates[0] if learning_rates else None ), "learning_rate_last": ( learning_rates[-1] if learning_rates else None ), "metrics": direct_metrics, "metric_slices": metric_slices, "parameter_count": parameter_count, "train_loss": history, "train_rows": int(len(train_rows)), "train_valid_targets_by_supervised_horizon": { str(horizon): int(count) for horizon, count in zip( horizons, train_target_mask[:, supervised_indices].sum(axis=0), strict=True, ) }, "updates": updates, "validation_rows": int(len(validation_rows)), "validation_valid_targets_by_horizon": { str(int(horizon)): int(count) for horizon, count in zip( EVALUATION_HORIZONS, validation_target_mask.sum(axis=0), strict=True, ) }, } del model if device.type == "mps": torch.mps.empty_cache() def fold_mean(slice_name: str) -> float: return float( np.mean( [ fold_results[name]["metric_slices"][slice_name][ "macro_nll" ] for name in DISCOVERY_FOLDS ] ) ) summary = { **config, "duration_seconds": time.monotonic() - started, "effective_batch_size": batch_size * gradient_accumulation_steps, "folds": fold_results, "mean_canonical_macro_nll": fold_mean("canonical"), "mean_dense_unseen_macro_nll": fold_mean("dense_unseen"), "mean_interpolation_baseline_dense_unseen_macro_nll": fold_mean( "interpolation_baseline_dense_unseen" ), "mean_prior_canonical_macro_nll": fold_mean("prior_canonical"), "mean_prior_dense_unseen_macro_nll": fold_mean( "prior_dense_unseen" ), "mean_supervised_macro_nll": fold_mean("supervised"), "run_manifest_sha256": run_manifest_sha256, } if withheld_indices: summary["mean_withheld_macro_nll"] = fold_mean("withheld") summary[ "mean_interpolation_baseline_withheld_macro_nll" ] = fold_mean("interpolation_baseline_withheld") results_path.write_text( json.dumps(summary, indent=2, sort_keys=True) + "\n" ) return summary def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--study-dir", type=Path, required=True) parser.add_argument("--asset", required=True) parser.add_argument("--model-name", choices=MODEL_NAMES, required=True) parser.add_argument( "--supervised-horizons", default="2,4,8,16,32", ) parser.add_argument("--context-length", type=int, default=128) parser.add_argument("--seed", type=int, default=0) parser.add_argument("--output-dir", type=Path, required=True) parser.add_argument("--epochs", type=int, default=16) parser.add_argument("--batch-size", type=int, default=256) parser.add_argument("--inference-batch-size", type=int, default=512) parser.add_argument("--target-parameters", type=int, default=25_000) parser.add_argument("--torch-threads", type=int, default=6) parser.add_argument("--max-train-rows", type=int, default=100_000) parser.add_argument( "--device-name", choices=("cpu", "mps"), default="mps", ) parser.add_argument("--gradient-accumulation-steps", type=int, default=16) parser.add_argument( "--optimizer-recipe", choices=OPTIMIZER_RECIPES, default="adamw_constant", ) parser.add_argument("--rotary-base", type=float, default=16.0) parser.add_argument( "--feature-set", choices=tuple(FEATURE_SETS), default="baseline", ) parser.add_argument("--diagnostic-tail-fraction", type=float, default=0.1) parser.add_argument("--early-stopping-patience", type=int, default=2) parser.add_argument("--early-stopping-min-delta", type=float, default=0.0005) args = parser.parse_args() print(json.dumps(run(**vars(args)), sort_keys=True)) if __name__ == "__main__": main()