Skip to content

Commit 779b893

Browse files
authored
Merge pull request #495 from woleigegg/fix/multifactor-lazy-calculation-thread-safe-pr
fix(multifactor): make lazy calculation thread-safe
2 parents 23ac116 + 65d58b9 commit 779b893

10 files changed

Lines changed: 883 additions & 147 deletions

File tree

hikyuu_cpp/hikyuu/trade_sys/multifactor/MultiFactorBase.cpp

Lines changed: 57 additions & 130 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,6 @@
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"
@@ -21,6 +20,7 @@
2120
#include "hikyuu/StockManager.h"
2221
#include "MultiFactorBase.h"
2322
#include "industry_neutralize.h"
23+
#include "StyleRegression.h"
2424

2525
namespace hku {
2626

@@ -142,7 +142,7 @@ void MultiFactorBase::baseCheckParam(const string& name) const {
142142
}
143143

144144
void MultiFactorBase::paramChanged() {
145-
m_calculated = false;
145+
m_calculated.store(false, std::memory_order_relaxed);
146146
}
147147

148148
void 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

181186
MultiFactorPtr 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

228233
void MultiFactorBase::setQuery(const KQuery& query) {
229234
m_query = query;
230-
m_calculated = false;
235+
m_calculated.store(false, std::memory_order_relaxed);
231236
}
232237

233238
void 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

242247
void 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

252257
void 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

258263
void MultiFactorBase::setNormalize(NormPtr norm) {
259264
m_norm = norm;
260-
m_calculated = false;
265+
m_calculated.store(false, std::memory_order_relaxed);
261266
}
262267

263268
void 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

301306
const DatetimeList& MultiFactorBase::getDatetimeList() {
302-
if (!m_calculated) {
303-
calculate();
304-
}
307+
calculate();
305308
return m_ref_dates;
306309
}
307310

308311
const 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

319320
const 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

328329
ScoreRecordList 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

437436
const 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

616532
vector<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

816729
void 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

Comments
 (0)