| """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 |
|
|