1919
2020import argparse
2121import csv
22+ import hashlib
2223import json
2324import math
2425import sys
@@ -88,6 +89,13 @@ def _parse_str_list(value: str) -> list[str]:
8889 return [part .strip () for part in str (value ).split ("," ) if part .strip ()]
8990
9091
92+ def _file_cache_tag (path : Path ) -> str :
93+ resolved = Path (path ).resolve ()
94+ stat = resolved .stat ()
95+ raw = f"{ resolved } :{ stat .st_size } :{ int (stat .st_mtime )} "
96+ return hashlib .sha1 (raw .encode ("utf-8" )).hexdigest ()[:12 ]
97+
98+
9199def _softmax (values : np .ndarray , * , temp : float = 1.0 ) -> np .ndarray :
92100 arr = np .asarray (values , dtype = np .float64 ) / max (1e-9 , float (temp ))
93101 arr = arr - float (np .max (arr ))
@@ -208,26 +216,53 @@ def _linear_channels(data: MktdData, rules: Sequence[LinearRule]) -> list[tuple[
208216 return out
209217
210218
211- def _xgb_channels (
219+ def _train_xgb_channel_models (
212220 train_data : MktdData ,
213- data : MktdData ,
214221 * ,
215222 experiment_names : Sequence [str ],
216223 rounds : int ,
217224 device : str ,
218- ) -> list [tuple [str , np .ndarray ]]:
225+ model_dir : Path | None = None ,
226+ cache_tag : str = "" ,
227+ ) -> list [tuple [str , object ]]:
219228 if not experiment_names :
220229 return []
221230 by_name = {exp .name : exp for exp in _experiments ()}
222231 missing = sorted (set (experiment_names ) - set (by_name ))
223232 if missing :
224233 raise ValueError (f"unknown XGB experiments: { ', ' .join (missing )} " )
225- out : list [tuple [str , np .ndarray ]] = []
226- valid = _valid_mask (data )
227- for name in experiment_names :
234+ out : list [tuple [str , object ]] = []
235+ if model_dir is not None :
236+ Path (model_dir ).mkdir (parents = True , exist_ok = True )
237+ for idx , name in enumerate (experiment_names , start = 1 ):
228238 exp = by_name [name ]
239+ cache_path = (
240+ Path (model_dir ) / f"{ name } _rounds{ int (rounds )} _{ str (device )} _{ cache_tag } .json"
241+ if model_dir is not None
242+ else None
243+ )
244+ if cache_path is not None and cache_path .exists ():
245+ print (f"xgb load { idx } /{ len (experiment_names )} { name } { cache_path } " , flush = True )
246+ import xgboost as xgb
247+
248+ model = xgb .Booster ()
249+ model .load_model (str (cache_path ))
250+ out .append ((name , model ))
251+ continue
252+ print (f"xgb train { idx } /{ len (experiment_names )} { name } horizon={ exp .horizon } label={ exp .label } " , flush = True )
229253 x_train , y_train = _build_dataset (train_data , horizon = exp .horizon , label = exp .label )
230254 model = _train_xgb (x_train , y_train , exp , rounds = int (rounds ), device = str (device ))
255+ if cache_path is not None :
256+ model .save_model (str (cache_path ))
257+ out .append ((name , model ))
258+ return out
259+
260+
261+ def _xgb_channels (data : MktdData , models : Sequence [tuple [str , object ]]) -> list [tuple [str , np .ndarray ]]:
262+ valid = _valid_mask (data )
263+ out : list [tuple [str , np .ndarray ]] = []
264+ for idx , (name , model ) in enumerate (models , start = 1 ):
265+ print (f"xgb score { idx } /{ len (models )} { name } T={ data .num_timesteps } " , flush = True )
231266 out .append ((f"xgb_in_sample:{ name } " , _normalize_score_matrix (_precompute_scores (data , model ), valid )))
232267 return out
233268
@@ -236,28 +271,15 @@ def _build_bank(
236271 data : MktdData ,
237272 * ,
238273 rules : Sequence [LinearRule ],
239- xgb_train_data : MktdData | None ,
240- xgb_experiment_names : Sequence [str ],
241- xgb_rounds : int ,
242- xgb_device : str ,
274+ xgb_models : Sequence [tuple [str , object ]],
243275 include_handcrafted : bool ,
244276) -> ScoreBank :
245277 channels : list [tuple [str , np .ndarray ]] = []
246278 channels .extend (_linear_channels (data , rules ))
247279 if include_handcrafted :
248280 channels .extend (_handcrafted_channels (data ))
249- if xgb_experiment_names :
250- if xgb_train_data is None :
251- raise ValueError ("xgb_train_data is required when XGB channels are requested" )
252- channels .extend (
253- _xgb_channels (
254- xgb_train_data ,
255- data ,
256- experiment_names = xgb_experiment_names ,
257- rounds = int (xgb_rounds ),
258- device = str (xgb_device ),
259- )
260- )
281+ if xgb_models :
282+ channels .extend (_xgb_channels (data , xgb_models ))
261283 names = [name for name , _scores in channels ]
262284 scores = np .stack ([scores for _name , scores in channels ], axis = 0 ).astype (np .float64 , copy = False )
263285 return ScoreBank (names = names , scores = scores )
@@ -603,6 +625,18 @@ def _sortino_from_equity(equity: np.ndarray, *, periods_per_year: float = 365.0)
603625 return float (returns .mean () / denom * np .sqrt (float (periods_per_year )))
604626
605627
628+ def _evolve_weights_after_return (target : np .ndarray , gross_return : np .ndarray , growth : float ) -> np .ndarray :
629+ if not np .isfinite (growth ) or float (growth ) <= 1e-8 :
630+ return np .zeros_like (np .asarray (target , dtype = np .float64 ), dtype = np .float64 )
631+ with np .errstate (divide = "ignore" , invalid = "ignore" , over = "ignore" ):
632+ weights = np .where (np .abs (target ) > 1e-12 , np .asarray (target , dtype = np .float64 ) * gross_return / growth , 0.0 )
633+ weights = np .where (np .isfinite (weights ), weights , 0.0 )
634+ gross = float (np .abs (weights ).sum ())
635+ if gross > 10.0 :
636+ weights *= 10.0 / gross
637+ return weights
638+
639+
606640def _simulate_vector_window (
607641 data : MktdData ,
608642 scores_by_t : np .ndarray ,
@@ -657,7 +691,7 @@ def _simulate_vector_window(
657691 growth = max (1e-9 , 1.0 + pnl - cost - borrow )
658692 equity = max (1e-9 , equity * growth )
659693 gross_return = np .where (valid , 1.0 + day_ret , 1.0 )
660- weights = np . where ( np . abs ( target ) > 1e-12 , target * gross_return / growth , 0.0 )
694+ weights = _evolve_weights_after_return ( target , gross_return , growth )
661695 curve .append (equity )
662696 final_turnover = float (np .abs (weights ).sum ())
663697 if final_turnover > 1e-9 :
@@ -1031,7 +1065,7 @@ def _simulate_vector_trace(
10311065 growth = max (1e-9 , 1.0 + pnl - cost - borrow )
10321066 equity = max (1e-9 , equity * growth )
10331067 gross_return = np .where (valid , 1.0 + day_ret , 1.0 )
1034- weights = np . where ( np . abs ( target ) > 1e-12 , target * gross_return / growth , 0.0 )
1068+ weights = _evolve_weights_after_return ( target , gross_return , growth )
10351069 curve .append (float (equity * initial_cash ))
10361070 positions_by_bar .append (
10371071 _position_rows_for_weights (
@@ -1474,6 +1508,7 @@ def main() -> int:
14741508 parser .add_argument ("--xgb-experiment-names" , default = "" )
14751509 parser .add_argument ("--xgb-rounds" , type = int , default = 80 )
14761510 parser .add_argument ("--xgb-device" , default = "cuda" )
1511+ parser .add_argument ("--xgb-model-dir" , type = Path , default = None )
14771512 parser .add_argument ("--no-handcrafted" , action = "store_true" )
14781513 parser .add_argument ("--out" , type = Path , default = Path ("analysis/binance33_meta_anneal.csv" ))
14791514 parser .add_argument ("--eval-days" , type = int , default = 120 )
@@ -1523,22 +1558,24 @@ def main() -> int:
15231558 max_rules = int (args .max_linear_rules ),
15241559 )
15251560 xgb_names = _parse_str_list (args .xgb_experiment_names )
1561+ xgb_models = _train_xgb_channel_models (
1562+ full_train_data ,
1563+ experiment_names = xgb_names ,
1564+ rounds = int (args .xgb_rounds ),
1565+ device = str (args .xgb_device ),
1566+ model_dir = args .xgb_model_dir ,
1567+ cache_tag = _file_cache_tag (args .train_data ) if xgb_names else "" ,
1568+ )
15261569 train_bank = _build_bank (
15271570 train_data ,
15281571 rules = rules ,
1529- xgb_train_data = full_train_data ,
1530- xgb_experiment_names = xgb_names ,
1531- xgb_rounds = int (args .xgb_rounds ),
1532- xgb_device = str (args .xgb_device ),
1572+ xgb_models = xgb_models ,
15331573 include_handcrafted = not bool (args .no_handcrafted ),
15341574 )
15351575 val_bank = _build_bank (
15361576 val_data ,
15371577 rules = rules ,
1538- xgb_train_data = full_train_data if xgb_names else None ,
1539- xgb_experiment_names = xgb_names ,
1540- xgb_rounds = int (args .xgb_rounds ),
1541- xgb_device = str (args .xgb_device ),
1578+ xgb_models = xgb_models ,
15421579 include_handcrafted = not bool (args .no_handcrafted ),
15431580 )
15441581 if train_bank .names != val_bank .names :
0 commit comments