66 */
77
88#include < cmath>
9- #include < Eigen/Dense>
109#include " hikyuu/utilities/thread/algorithm.h"
1110#include " hikyuu/indicator/crt/ALIGN.h"
1211#include " hikyuu/indicator/crt/KDATA.h"
2120#include " hikyuu/StockManager.h"
2221#include " MultiFactorBase.h"
2322#include " industry_neutralize.h"
23+ #include " StyleRegression.h"
2424
2525namespace hku {
2626
@@ -142,7 +142,7 @@ void MultiFactorBase::baseCheckParam(const string& name) const {
142142}
143143
144144void MultiFactorBase::paramChanged () {
145- m_calculated = false ;
145+ m_calculated. store ( false , std::memory_order_relaxed) ;
146146}
147147
148148void MultiFactorBase::_checkData () {
@@ -165,17 +165,22 @@ void MultiFactorBase::_checkData() {
165165 m_stks.size ());
166166}
167167
168- void MultiFactorBase::reset () {
169- _reset ();
170-
171- std::lock_guard<std::mutex> lock (m_mutex);
168+ void MultiFactorBase::clearCalculatedData () {
172169 m_ref_dates = {};
173170 m_stk_map = {};
174171 m_all_factors = {};
175172 m_date_index = {};
176173 m_stk_factor_by_date = {};
177174 m_ic = {};
178- m_calculated = false ;
175+ }
176+
177+ void MultiFactorBase::reset () {
178+ // 全程持锁:避免与正在进行的 calculate 写写交叉。
179+ // 注意:_reset() 为虚函数,自定义实现不得在锁内重入同一实例需要 m_mutex 的方法。
180+ std::lock_guard<std::mutex> lock (m_mutex);
181+ _reset ();
182+ clearCalculatedData ();
183+ m_calculated.store (false , std::memory_order_release);
179184}
180185
181186MultiFactorPtr MultiFactorBase::clone () {
@@ -212,7 +217,7 @@ MultiFactorPtr MultiFactorBase::clone() {
212217
213218 p->m_special_category = m_special_category;
214219
215- p->m_calculated = false ;
220+ p->m_calculated . store ( false , std::memory_order_relaxed) ;
216221 // 强制重算,不克隆以下缓存,避免非线程安全
217222 // p->m_stk_map = m_stk_map;
218223 // p->m_date_index = m_date_index;
@@ -227,7 +232,7 @@ MultiFactorPtr MultiFactorBase::clone() {
227232
228233void MultiFactorBase::setQuery (const KQuery& query) {
229234 m_query = query;
230- m_calculated = false ;
235+ m_calculated. store ( false , std::memory_order_relaxed) ;
231236}
232237
233238void MultiFactorBase::setRefStock (const Stock& stk) {
@@ -236,7 +241,7 @@ void MultiFactorBase::setRefStock(const Stock& stk) {
236241 HKU_CHECK (ref_dates.size () >= 2 , " The dates len is insufficient! current len: {}" ,
237242 ref_dates.size ());
238243 m_ref_stk = tmp_stk;
239- m_calculated = false ;
244+ m_calculated. store ( false , std::memory_order_relaxed) ;
240245}
241246
242247void MultiFactorBase::setStockList (const StockList& stks) {
@@ -246,18 +251,18 @@ void MultiFactorBase::setStockList(const StockList& stks) {
246251 }
247252
248253 m_stks = stks;
249- m_calculated = false ;
254+ m_calculated. store ( false , std::memory_order_relaxed) ;
250255}
251256
252257void MultiFactorBase::setRefFactorSet (const FactorSet& factorset) {
253258 HKU_CHECK (!factorset.isNull () && !factorset.empty (), " Input factor set is null or empty!" );
254259 m_factorset = factorset;
255- m_calculated = false ;
260+ m_calculated. store ( false , std::memory_order_relaxed) ;
256261}
257262
258263void MultiFactorBase::setNormalize (NormPtr norm) {
259264 m_norm = norm;
260- m_calculated = false ;
265+ m_calculated. store ( false , std::memory_order_relaxed) ;
261266}
262267
263268void MultiFactorBase::addSpecialNormalize (const string& name, NormalizePtr norm,
@@ -295,29 +300,25 @@ void MultiFactorBase::addSpecialNormalize(const string& name, NormalizePtr norm,
295300 m_special_style_inds[found_name] = style_inds;
296301 }
297302
298- m_calculated = false ;
303+ m_calculated. store ( false , std::memory_order_relaxed) ;
299304}
300305
301306const DatetimeList& MultiFactorBase::getDatetimeList () {
302- if (!m_calculated) {
303- calculate ();
304- }
307+ calculate ();
305308 return m_ref_dates;
306309}
307310
308311const Indicator& MultiFactorBase::getFactor (const Stock& stk) {
309312 HKU_CHECK (getParam<bool >(" save_all_factors" ),
310313 " param \" save_all_factors\" is false, can't get all factors!" );
311- if (!m_calculated) {
312- calculate ();
313- }
314+ calculate ();
314315 const auto iter = m_stk_map.find (stk);
315316 HKU_CHECK (iter != m_stk_map.cend (), " Could not find this stock: {}" , stk);
316317 return m_all_factors[iter->second ];
317318}
318319
319320const IndicatorList& MultiFactorBase::getAllFactors () {
320- if (getParam<bool >(" save_all_factors" ) && !m_calculated ) {
321+ if (getParam<bool >(" save_all_factors" )) {
321322 calculate ();
322323 } else {
323324 HKU_WARN (" param \" save_all_factors\" is false, can't get all factors!" );
@@ -326,9 +327,7 @@ const IndicatorList& MultiFactorBase::getAllFactors() {
326327}
327328
328329ScoreRecordList MultiFactorBase::getScores (const Datetime& d) {
329- if (!m_calculated) {
330- calculate ();
331- }
330+ calculate ();
332331 ScoreRecordList ret;
333332 const auto iter = m_date_index.find (d);
334333 HKU_IF_RETURN (iter == m_date_index.cend (), ret);
@@ -435,9 +434,7 @@ ScoreRecordList MultiFactorBase::getScores(const Datetime& date, size_t start, s
435434}
436435
437436const vector<ScoreRecordList>& MultiFactorBase::getAllScores () {
438- if (!m_calculated) {
439- calculate ();
440- }
437+ calculate ();
441438 return m_stk_factor_by_date;
442439}
443440
@@ -446,9 +443,7 @@ Indicator MultiFactorBase::getIC(int ndays) {
446443 htr (" mf param \" save_all_factors\" is false, can't get all factors!, please "
447444 " set it to true if you want to get IC/ICIR!" ));
448445
449- if (!m_calculated) {
450- calculate ();
451- }
446+ calculate ();
452447
453448 std::lock_guard<std::mutex> lock (m_mutex);
454449
@@ -531,87 +526,8 @@ IndicatorList MultiFactorBase::_getAllReturns(int ndays) const {
531526
532527// 行业中性化(按行业分组去组内均值)的纯函数实现见 industry_neutralize.h,
533528// 提取为内部 inline header 供白盒单元测试直接包含调用。
534-
535- // 计算多元回归中性化后的因子,y为因子,x为多个解释变量(包含常数项)- Eigen版本
536- static PriceList calculate_residuals (const PriceList& y, const std::vector<PriceList>& x) {
537- HKU_ASSERT (!x.empty ());
538- size_t n = y.size ();
539- for (const auto & xi : x) {
540- HKU_ASSERT (xi.size () == n);
541- }
542-
543- PriceList residuals (n, Null<price_t >());
544- size_t k = x.size (); // 解释变量个数
545-
546- // 构建设计矩阵和因变量向量
547- Eigen::MatrixXd Xmat (n, k + 1 );
548- Eigen::VectorXd Yvec (n);
549-
550- // 填充数据 - 第一列为常数项(全1)
551- Xmat.col (0 ).setConstant (1.0 );
552-
553- // 标记有效数据点
554- std::vector<bool > valid (n, true );
555-
556- for (size_t i = 0 ; i < n; ++i) {
557- Yvec (i) = y[i];
558-
559- // 检查因变量是否有效
560- if (std::isnan (y[i]) || std::isinf (y[i])) {
561- valid[i] = false ;
562- continue ;
563- }
564-
565- // 填充自变量并检查有效性
566- for (size_t j = 0 ; j < k; ++j) {
567- Xmat (i, j + 1 ) = x[j][i];
568- if (std::isnan (x[j][i]) || std::isinf (x[j][i])) {
569- valid[i] = false ;
570- break ;
571- }
572- }
573- }
574-
575- // 计算有效数据点数量
576- size_t valid_count = std::count (valid.begin (), valid.end (), true );
577-
578- // 数据点不足
579- if (valid_count <= k + 1 ) {
580- return residuals;
581- }
582-
583- // 创建有效数据的子矩阵
584- Eigen::MatrixXd X_valid (valid_count, k + 1 );
585- Eigen::VectorXd Y_valid (valid_count);
586-
587- size_t valid_idx = 0 ;
588- for (size_t i = 0 ; i < n; ++i) {
589- if (valid[i]) {
590- X_valid.row (valid_idx) = Xmat.row (i);
591- Y_valid (valid_idx) = Yvec (i);
592- valid_idx++;
593- }
594- }
595-
596- // 使用QR分解求解线性回归 β = (X'X)^(-1)X'Y
597- Eigen::VectorXd beta = X_valid.colPivHouseholderQr ().solve (Y_valid);
598-
599- // 检查解是否有效
600- if (beta.hasNaN ()) {
601- return residuals;
602- }
603-
604- // 计算拟合值和残差
605- Eigen::VectorXd fitted = Xmat * beta;
606-
607- for (size_t i = 0 ; i < n; ++i) {
608- if (valid[i]) {
609- residuals[i] = y[i] - fitted (i);
610- }
611- }
612-
613- return residuals;
614- }
529+ // 风格因子中性化残差回归实现见 StyleRegression.cpp,从本类中提取为串行内核,
530+ // 不再在运行时修改进程级 Eigen 线程配置。
615531
616532vector<IndicatorList> MultiFactorBase::getAllSrcFactors () {
617533 vector<IndicatorList> all_stk_inds;
@@ -675,9 +591,9 @@ vector<IndicatorList> MultiFactorBase::getAllSrcFactors() {
675591
676592 // 时间截面标准化/归一化
677593 if (m_norm || !m_special_category.empty () || !m_special_style_inds.empty ()) {
678- // 压制 Eigen 内部 OpenMP 并行,避免与外层按日线程池叠加导致线程超载;
679- // calculate_residuals 内的 Eigen 矩阵均为栈局部对象,外层按日并行天然可重入。
680- Eigen::setNbThreads ( 1 );
594+ // 风格因子中性化残差回归已提取为串行内核(StyleRegression.cpp),
595+ // 不再在运行时修改进程级 Eigen::setNbThreads,避免并发 MF 互相污染全局配置;
596+ // 外层按日并行天然可重入,回归内部均为栈局部对象。
681597 unordered_map<string, std::pair<PriceList, size_t >> ind_dummy_dict = _buildDummyIndex ();
682598 global_parallel_for_index_void (
683599 0 , days_total,
@@ -735,7 +651,7 @@ vector<IndicatorList> MultiFactorBase::getAllSrcFactors() {
735651 style_value[si] = per_factor[j][si][di];
736652 }
737653 }
738- new_value = calculate_residuals (new_value, style_value_day);
654+ new_value = calculate_style_residuals (new_value, style_value_day);
739655 }
740656
741657 for (size_t si = 0 ; si < stk_count; si++) {
@@ -744,9 +660,6 @@ vector<IndicatorList> MultiFactorBase::getAllSrcFactors() {
744660 }
745661 }
746662 });
747-
748- // 恢复 Eigen 线程数
749- Eigen::setNbThreads (std::thread::hardware_concurrency ());
750663 }
751664
752665 return all_stk_inds;
@@ -814,12 +727,24 @@ void MultiFactorBase::_buildIndex() {
814727}
815728
816729void MultiFactorBase::calculate () {
817- HKU_IF_RETURN (m_calculated, void ());
730+ // Fast path: lock-free acquire 检查是否已 Ready
731+ if (m_calculated.load (std::memory_order_acquire)) {
732+ return ;
733+ }
818734
819735 std::lock_guard<std::mutex> lock (m_mutex);
820- _checkData ();
736+
737+ // 锁内二次检查:mutex 已提供慢路径同步,relaxed 即可
738+ if (m_calculated.load (std::memory_order_relaxed)) {
739+ return ;
740+ }
741+
742+ // 构建前清理旧结果,确保重试基于干净状态
743+ clearCalculatedData ();
821744
822745 try {
746+ _checkData ();
747+
823748 { // 获取所有证券所有对齐后的原始因子
824749 vector<IndicatorList> all_stk_inds = getAllSrcFactors ();
825750
@@ -839,19 +764,21 @@ void MultiFactorBase::calculate() {
839764
840765 // 计算完成后创建截面索引
841766 _buildIndex ();
842- } catch (const std::exception& e) {
843- HKU_ERROR (e.what ());
844- } catch (...) {
845- HKU_ERROR_UNKNOWN ;
846- }
847767
848- if (!getParam<bool >(" save_all_factors" )) {
849- m_all_factors = {};
850- m_stk_map = {};
768+ if (!getParam<bool >(" save_all_factors" )) {
769+ m_all_factors = {};
770+ m_stk_map = {};
771+ }
772+ } catch (...) {
773+ // 失败清理:所有异步子任务已在 wait_and_drain 语义下结束,
774+ // 清除基类半成品,保持未计算状态,原异常向上传播,允许下一调用者重试。
775+ clearCalculatedData ();
776+ m_calculated.store (false , std::memory_order_relaxed);
777+ throw ;
851778 }
852779
853- // 更新计算状态
854- m_calculated = true ;
780+ // Publish:release 保证此前所有写入对后续 acquire 读取可见
781+ m_calculated. store ( true , std::memory_order_release) ;
855782}
856783
857784} // namespace hku
0 commit comments