quant_test / factor_engine /gp /fitness.py
lucky-loster's picture
Upload folder using huggingface_hub
590a501 verified
Raw
History Blame Contribute Delete
5.93 kB
"""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