File size: 5,930 Bytes
590a501
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
"""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