"""Fitness evaluation with anti-overfit and diversity penalties.""" from __future__ import annotations from dataclasses import dataclass import numpy as np import torch from factor_engine.gp.operators import RISKY_OPERATOR_NAMES @dataclass class FitnessConfig: min_stocks: int = 50 min_tree_nodes: int = 3 depth_penalty: float = 0.006 node_penalty: float = 0.0012 duplicate_penalty: float = 0.05 elite_corr_fatal: float = 0.95 elite_corr_hard: float = 0.85 elite_corr_mid: float = 0.75 elite_corr_soft: float = 0.65 export_corr_threshold: float = 0.80 clip_near_value: float = 999999.0 huge_abs_value: float = 1e5 min_valid_factor_ratio: float = 0.50 min_unique_factor_values: int = 20 def calculate_rank_ic_series(factor_tensor, target_tensor, mask=None, min_stocks=5): valid_mask = ~torch.isnan(factor_tensor) & ~torch.isnan(target_tensor) if mask is not None: valid_mask = valid_mask & mask factor = torch.where(valid_mask, factor_tensor, torch.zeros_like(factor_tensor)) target = torch.where(valid_mask, target_tensor, torch.zeros_like(target_tensor)) f_rank = factor.argsort(dim=0).argsort(dim=0).float() r_rank = target.argsort(dim=0).argsort(dim=0).float() count = valid_mask.sum(dim=0).float() mu_f = torch.where(count > 1, (f_rank * valid_mask.float()).sum(dim=0) / count, torch.zeros_like(count)) mu_r = torch.where(count > 1, (r_rank * valid_mask.float()).sum(dim=0) / count, torch.zeros_like(count)) df = (f_rank - mu_f) * valid_mask.float() dr = (r_rank - mu_r) * valid_mask.float() num = (df * dr).sum(dim=0) den = torch.sqrt((df ** 2).sum(dim=0)) * torch.sqrt((dr ** 2).sum(dim=0)) return torch.where( (den > 1e-8) & (count >= min_stocks), num / den, torch.full_like(den, float("nan")), ) def factor_report(factor, target, mask=None, min_stocks=50): ic = calculate_rank_ic_series(factor, target, mask=mask, min_stocks=min_stocks) valid_ic = ic[~torch.isnan(ic)] if valid_ic.numel() < 20: return {"ic_mean": np.nan, "ic_std": np.nan, "icir": np.nan, "pos_ratio": np.nan, "n_ic": int(valid_ic.numel())} ic_mean = valid_ic.mean().item() ic_std = valid_ic.std(unbiased=False).item() return { "ic_mean": ic_mean, "ic_std": ic_std, "icir": ic_mean / (ic_std + 1e-6), "pos_ratio": (valid_ic > 0).float().mean().item(), "n_ic": int(valid_ic.numel()), } def sampled_spearman_corr_torch(x, y, mask, max_points=100_000): valid = ~torch.isnan(x) & ~torch.isnan(y) if mask is not None: valid = valid & mask idx = valid.flatten().nonzero(as_tuple=False).flatten() if idx.numel() < 1000: return np.nan if idx.numel() > max_points: idx = idx[torch.randperm(idx.numel(), device=idx.device)[:max_points]] xf, yf = x.flatten()[idx], y.flatten()[idx] xr, yr = xf.argsort().argsort().float(), yf.argsort().argsort().float() xr, yr = xr - xr.mean(), yr - yr.mean() denom = torch.sqrt((xr ** 2).sum()) * torch.sqrt((yr ** 2).sum()) return np.nan if denom <= 1e-8 else float((xr * yr).sum().item() / denom.item()) def operator_repetition_penalty(tree): counts = {} for node in tree.get_nodes(): if hasattr(node, "name"): base = node.name.split("_")[0] counts[base] = counts.get(base, 0) + 1 penalty = sum(0.04 * (cnt - 3) for name, cnt in counts.items() if cnt >= 4) penalty += sum(0.08 * (counts.get(name, 0) - 2) for name in RISKY_OPERATOR_NAMES if counts.get(name, 0) >= 3) return float(penalty) def factor_extreme_penalty(factor, mask=None, cfg: FitnessConfig | None = None): cfg = cfg or FitnessConfig() valid = ~torch.isnan(factor) if mask is not None: valid = valid & mask x = factor[valid] if x.numel() < 1000: return 0.50 if torch.unique(x[: min(x.numel(), 100_000)]).numel() < cfg.min_unique_factor_values: return 0.50 clip_ratio = (torch.abs(x) >= cfg.clip_near_value).float().mean().item() huge_ratio = (torch.abs(x) >= cfg.huge_abs_value).float().mean().item() return min(0.50, clip_ratio * 100) + min(0.40, huge_ratio * 40) def calculate_fitness( tree, factor, target, train_mask, test_mask, cfg: FitnessConfig, formula_seen=None, elite_factor_cache=None, ): is_report = factor_report(factor, target, mask=train_mask, min_stocks=cfg.min_stocks) oos_report = factor_report(factor, target, mask=test_mask, min_stocks=cfg.min_stocks) if np.isnan(is_report["icir"]) or tree.get_size() < cfg.min_tree_nodes: return -999.0, is_report, oos_report duplicate_penalty = cfg.duplicate_penalty if formula_seen and str(tree) in formula_seen else 0.0 corr_penalty, fatal_duplicate = 0.0, False if elite_factor_cache: corrs = [ abs(sampled_spearman_corr_torch(factor, old, mask=train_mask, max_points=50_000)) for old in elite_factor_cache ] valid_corrs = [c for c in corrs if not np.isnan(c)] if valid_corrs: m = max(valid_corrs) if m >= cfg.elite_corr_fatal: fatal_duplicate = True elif m >= cfg.elite_corr_hard: corr_penalty = 0.50 elif m >= cfg.elite_corr_mid: corr_penalty = 0.25 elif m >= cfg.elite_corr_soft: corr_penalty = 0.10 if fatal_duplicate: return -999.0, is_report, oos_report complexity = tree.get_depth() * cfg.depth_penalty + tree.get_size() * cfg.node_penalty fitness = ( is_report["icir"] - complexity - duplicate_penalty - corr_penalty - operator_repetition_penalty(tree) - factor_extreme_penalty(factor, train_mask, cfg) ) return float(fitness), is_report, oos_report