A new strategy is added haha
This commit is contained in:
+14
-8
@@ -62,15 +62,15 @@ class ChanLunClassifier:
|
||||
# 默认XGBoost参数
|
||||
default_params = {
|
||||
'objective': 'binary:logistic',
|
||||
'max_depth': 4,
|
||||
'eta': 0.03,
|
||||
'max_depth': 8,
|
||||
'eta': 0.01,
|
||||
'subsample': 0.8,
|
||||
'colsample_bytree': 0.8,
|
||||
'eval_metric': 'auc',
|
||||
'gamma': 0.1,
|
||||
'min_child_weight': 3,
|
||||
'alpha': 1, # L1正则化
|
||||
'lambda': 3, # L2正则化
|
||||
'gamma': 0.0,
|
||||
'min_child_weight': 1,
|
||||
'alpha': 0, # L1正则化
|
||||
'lambda': 0.5, # L2正则化
|
||||
'scale_pos_weight': 1
|
||||
}
|
||||
|
||||
@@ -208,7 +208,7 @@ class ChanLunClassifier:
|
||||
else:
|
||||
self.model = xgb.Booster()
|
||||
self.model.load_model(model_file_path)
|
||||
def find_best_params(self, dataframe=None, save_csv=False, csv_path_prefix='param_'):
|
||||
def find_best_params(self, dataframe=None, save_csv=False, csv_path_prefix='param_', model_name=None):
|
||||
"""
|
||||
寻找最佳参数组合
|
||||
:param dataframe: 输入的DataFrame,如果为None则使用初始化时的dataframe
|
||||
@@ -240,7 +240,7 @@ class ChanLunClassifier:
|
||||
param_info = f"eta{params['eta']}_depth{params['max_depth']}"
|
||||
train_csv_path = f"{csv_path_prefix}train_{param_info}.csv" if save_csv else None
|
||||
|
||||
model = self.train_model(dataframe=dataframe, data_file_path=train_csv_path, custom_params=params)
|
||||
model = self.train_model(dataframe=dataframe, data_file_path=train_csv_path, custom_params=params, model_name=model_name)
|
||||
|
||||
# 分割数据集,后20%用于测试
|
||||
if dataframe is None:
|
||||
@@ -298,6 +298,8 @@ class ChanLunClassifier:
|
||||
for klc in klc_list:
|
||||
if klc.klc_fx_type != Chan_KLC_FX.UNKNOWN:
|
||||
sample_list.append(klc)
|
||||
klc_count = 0
|
||||
print('Processing data...')
|
||||
for klc in sample_list:
|
||||
if bi_index >= len(bi_list):
|
||||
bi_index = len(bi_list) - 1
|
||||
@@ -333,6 +335,10 @@ class ChanLunClassifier:
|
||||
|
||||
feature_data.append(feature_vec)
|
||||
labels.append(label)
|
||||
klc_count += 1
|
||||
percent = klc_count/len(sample_list)*100
|
||||
if percent % 10 == 0:
|
||||
print('Data processed:', percent, '%')
|
||||
for index, key in enumerate(feature_keys):
|
||||
print(index, key, feature_data[0][index])
|
||||
# 如果需要保存到CSV
|
||||
|
||||
Reference in New Issue
Block a user