预编译数值表达式助手参数
This commit is contained in:
@@ -942,6 +942,49 @@ enum RuntimeExpressionSegment {
|
||||
struct RuntimeHelperBinding {
|
||||
name: String,
|
||||
args: Vec<String>,
|
||||
compiled_args: Option<CompiledRuntimeHelperArgs>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
enum CompiledRuntimeHelperArgs {
|
||||
RollingMean {
|
||||
field: String,
|
||||
lookback: usize,
|
||||
current: bool,
|
||||
},
|
||||
RollingMaxCurrent {
|
||||
field: String,
|
||||
lookback: usize,
|
||||
},
|
||||
RollingReturnStddevCurrent {
|
||||
field: String,
|
||||
return_count: usize,
|
||||
},
|
||||
VolumeMean {
|
||||
lookback: usize,
|
||||
},
|
||||
RollingAggregate {
|
||||
field: String,
|
||||
lookback: usize,
|
||||
operation: CompiledRollingOperation,
|
||||
},
|
||||
PctChange {
|
||||
field: String,
|
||||
lookback: usize,
|
||||
},
|
||||
FactorValue {
|
||||
field: String,
|
||||
lookback: usize,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum CompiledRollingOperation {
|
||||
Sum,
|
||||
Min,
|
||||
Max,
|
||||
Stddev,
|
||||
Zscore,
|
||||
}
|
||||
|
||||
/// Typed result of resolving a runtime helper.
|
||||
@@ -4894,8 +4937,11 @@ impl PlatformExprStrategy {
|
||||
binding: &RuntimeHelperBinding,
|
||||
expected_type: NumericVmValueType,
|
||||
) -> Result<NumericVmValue, BacktestError> {
|
||||
let resolved =
|
||||
self.resolve_runtime_helper(ctx, day, stock, &binding.name, &binding.args)?;
|
||||
let resolved = if let Some(compiled_args) = binding.compiled_args.as_ref() {
|
||||
self.resolve_compiled_runtime_helper(ctx, day, stock, &binding.name, compiled_args)?
|
||||
} else {
|
||||
self.resolve_runtime_helper(ctx, day, stock, &binding.name, &binding.args)?
|
||||
};
|
||||
match (expected_type, resolved) {
|
||||
(NumericVmValueType::Number, RuntimeHelperResolution::Number(value)) => {
|
||||
Ok(NumericVmValue::Number(value))
|
||||
@@ -4931,6 +4977,112 @@ impl PlatformExprStrategy {
|
||||
}
|
||||
}
|
||||
|
||||
fn resolve_compiled_runtime_helper(
|
||||
&self,
|
||||
ctx: &StrategyContext<'_>,
|
||||
day: &DayExpressionState,
|
||||
stock: Option<&StockExpressionState>,
|
||||
helper: &str,
|
||||
args: &CompiledRuntimeHelperArgs,
|
||||
) -> Result<RuntimeHelperResolution, BacktestError> {
|
||||
match args {
|
||||
CompiledRuntimeHelperArgs::RollingMean {
|
||||
field,
|
||||
lookback,
|
||||
current,
|
||||
} => {
|
||||
let value = if *current {
|
||||
self.resolve_current_rolling_mean(ctx, day, stock, field, *lookback)?
|
||||
} else {
|
||||
self.resolve_rolling_mean(ctx, day, stock, field, *lookback)?
|
||||
};
|
||||
Ok(RuntimeHelperResolution::Number(value))
|
||||
}
|
||||
CompiledRuntimeHelperArgs::RollingMaxCurrent { field, lookback } => {
|
||||
let values =
|
||||
self.resolve_current_rolling_values(ctx, day, stock, field, *lookback)?;
|
||||
let value = values.iter().copied().fold(f64::NEG_INFINITY, f64::max);
|
||||
Ok(Self::normalized_runtime_number(value))
|
||||
}
|
||||
CompiledRuntimeHelperArgs::RollingReturnStddevCurrent {
|
||||
field,
|
||||
return_count,
|
||||
} => {
|
||||
let values = self.resolve_current_rolling_values(
|
||||
ctx,
|
||||
day,
|
||||
stock,
|
||||
field,
|
||||
return_count.saturating_add(1),
|
||||
)?;
|
||||
let returns = values
|
||||
.windows(2)
|
||||
.map(|pair| pair[1] / pair[0] - 1.0)
|
||||
.collect::<Vec<_>>();
|
||||
if returns.iter().any(|value| !value.is_finite()) {
|
||||
return Err(BacktestError::Execution(format!(
|
||||
"invalid current rolling return for field {field} with count {return_count}"
|
||||
)));
|
||||
}
|
||||
Ok(Self::normalized_runtime_number(rolling_sample_stddev(
|
||||
&returns,
|
||||
)))
|
||||
}
|
||||
CompiledRuntimeHelperArgs::VolumeMean { lookback } => {
|
||||
let value = self.resolve_rolling_mean(ctx, day, stock, "volume", *lookback)?;
|
||||
Ok(RuntimeHelperResolution::Number(value))
|
||||
}
|
||||
CompiledRuntimeHelperArgs::RollingAggregate {
|
||||
field,
|
||||
lookback,
|
||||
operation,
|
||||
} => {
|
||||
let values = self.resolve_rolling_values(ctx, day, stock, field, *lookback)?;
|
||||
let value = match operation {
|
||||
CompiledRollingOperation::Sum => values.iter().sum::<f64>(),
|
||||
CompiledRollingOperation::Min => {
|
||||
values.iter().copied().fold(f64::INFINITY, f64::min)
|
||||
}
|
||||
CompiledRollingOperation::Max => {
|
||||
values.iter().copied().fold(f64::NEG_INFINITY, f64::max)
|
||||
}
|
||||
CompiledRollingOperation::Stddev => rolling_stddev(&values),
|
||||
CompiledRollingOperation::Zscore => rolling_zscore(&values),
|
||||
};
|
||||
Ok(Self::normalized_runtime_number(value))
|
||||
}
|
||||
CompiledRuntimeHelperArgs::PctChange { field, lookback } => {
|
||||
let values = self.resolve_rolling_values(
|
||||
ctx,
|
||||
day,
|
||||
stock,
|
||||
field,
|
||||
lookback.saturating_add(1),
|
||||
)?;
|
||||
let first = values.first().copied().unwrap_or_default();
|
||||
let last = values.last().copied().unwrap_or_default();
|
||||
let value = if first.abs() <= f64::EPSILON {
|
||||
0.0
|
||||
} else {
|
||||
last / first - 1.0
|
||||
};
|
||||
Ok(Self::normalized_runtime_number(value))
|
||||
}
|
||||
CompiledRuntimeHelperArgs::FactorValue { field, lookback } => {
|
||||
let stock = stock.ok_or_else(|| {
|
||||
BacktestError::Execution(format!("{helper} requires stock context"))
|
||||
})?;
|
||||
let start = self.helper_start_date(ctx, day.date, *lookback);
|
||||
let value = ctx
|
||||
.get_factor(&stock.symbol, start, day.date, field)
|
||||
.last()
|
||||
.map(|row| row.value)
|
||||
.unwrap_or(0.0);
|
||||
Ok(Self::normalized_runtime_number(value))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn numeric_vm_identifier_value(
|
||||
&self,
|
||||
ctx: &StrategyContext<'_>,
|
||||
@@ -5888,6 +6040,7 @@ impl PlatformExprStrategy {
|
||||
let binding = RuntimeHelperBinding {
|
||||
name: name.clone(),
|
||||
args: args.clone(),
|
||||
compiled_args: Self::compile_runtime_helper_args(name, args),
|
||||
};
|
||||
if let Some(existing) = bindings.get(scope_name)
|
||||
&& (existing.name != binding.name || existing.args != binding.args)
|
||||
@@ -5902,6 +6055,81 @@ impl PlatformExprStrategy {
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
fn compile_runtime_helper_args(
|
||||
helper: &str,
|
||||
args: &[String],
|
||||
) -> Option<CompiledRuntimeHelperArgs> {
|
||||
let field_lookback = || {
|
||||
if args.len() != 2 {
|
||||
return None;
|
||||
}
|
||||
Some((
|
||||
Self::parse_string_or_identifier(&args[0]).ok()?,
|
||||
Self::parse_positive_usize(&args[1]).ok()?,
|
||||
))
|
||||
};
|
||||
match helper {
|
||||
"rolling_mean" | "sma" | "ma" => {
|
||||
let (field, lookback) = field_lookback()?;
|
||||
Some(CompiledRuntimeHelperArgs::RollingMean {
|
||||
field,
|
||||
lookback,
|
||||
current: false,
|
||||
})
|
||||
}
|
||||
"rolling_mean_current" => {
|
||||
let (field, lookback) = field_lookback()?;
|
||||
Some(CompiledRuntimeHelperArgs::RollingMean {
|
||||
field,
|
||||
lookback,
|
||||
current: true,
|
||||
})
|
||||
}
|
||||
"rolling_max_current" => {
|
||||
let (field, lookback) = field_lookback()?;
|
||||
Some(CompiledRuntimeHelperArgs::RollingMaxCurrent { field, lookback })
|
||||
}
|
||||
"rolling_return_stddev_current" => {
|
||||
let (field, return_count) = field_lookback()?;
|
||||
Some(CompiledRuntimeHelperArgs::RollingReturnStddevCurrent {
|
||||
field,
|
||||
return_count,
|
||||
})
|
||||
}
|
||||
"vma" if args.len() == 1 => Some(CompiledRuntimeHelperArgs::VolumeMean {
|
||||
lookback: Self::parse_positive_usize(&args[0]).ok()?,
|
||||
}),
|
||||
"rolling_sum" | "rolling_min" | "rolling_max" | "rolling_stddev" | "stddev"
|
||||
| "rolling_zscore" => {
|
||||
let (field, lookback) = field_lookback()?;
|
||||
let operation = match helper {
|
||||
"rolling_sum" => CompiledRollingOperation::Sum,
|
||||
"rolling_min" => CompiledRollingOperation::Min,
|
||||
"rolling_max" => CompiledRollingOperation::Max,
|
||||
"rolling_stddev" | "stddev" => CompiledRollingOperation::Stddev,
|
||||
"rolling_zscore" => CompiledRollingOperation::Zscore,
|
||||
_ => return None,
|
||||
};
|
||||
Some(CompiledRuntimeHelperArgs::RollingAggregate {
|
||||
field,
|
||||
lookback,
|
||||
operation,
|
||||
})
|
||||
}
|
||||
"pct_change" => {
|
||||
let (field, lookback) = field_lookback()?;
|
||||
Some(CompiledRuntimeHelperArgs::PctChange { field, lookback })
|
||||
}
|
||||
"factor_value" | "get_factor_value" if (1..=2).contains(&args.len()) => {
|
||||
Some(CompiledRuntimeHelperArgs::FactorValue {
|
||||
field: Self::parse_string_or_identifier(&args[0]).ok()?,
|
||||
lookback: Self::parse_optional_positive_usize(args.get(1), 1).ok()?,
|
||||
})
|
||||
}
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn numeric_vm_helper_type(helper: &str) -> Option<NumericVmValueType> {
|
||||
match helper {
|
||||
"has_dividend" | "has_split" | "is_margin_stock" => Some(NumericVmValueType::Boolean),
|
||||
@@ -6043,6 +6271,15 @@ impl PlatformExprStrategy {
|
||||
helper: &str,
|
||||
args: &[String],
|
||||
) -> Result<RuntimeHelperResolution, BacktestError> {
|
||||
if let Some(compiled_args) = Self::compile_runtime_helper_args(helper, args) {
|
||||
return self.resolve_compiled_runtime_helper(
|
||||
ctx,
|
||||
day,
|
||||
stock,
|
||||
helper,
|
||||
&compiled_args,
|
||||
);
|
||||
}
|
||||
match helper {
|
||||
"factor" => {
|
||||
let key = Self::normalize_runtime_factor_key(&Self::parse_string_or_identifier(
|
||||
@@ -12222,8 +12459,9 @@ mod tests {
|
||||
use chrono::{NaiveDate, NaiveTime};
|
||||
|
||||
use super::{
|
||||
PlatformAccountActionKind, PlatformExplicitActionStage, PlatformExplicitCancelKind,
|
||||
PlatformExplicitOrderKind, PlatformExprStrategy, PlatformExprStrategyConfig,
|
||||
CompiledRuntimeHelperArgs, PlatformAccountActionKind, PlatformExplicitActionStage,
|
||||
PlatformExplicitCancelKind, PlatformExplicitOrderKind, PlatformExprStrategy,
|
||||
PlatformExprStrategyConfig,
|
||||
PlatformPortfolioDrawdownControlConfig, PlatformPortfolioDrawdownController,
|
||||
PlatformRebalanceSchedule, PlatformScheduleFrequency, PlatformStopTakeReferencePriceMode,
|
||||
PlatformTradeAction, PlatformUniverseActionKind, RuntimeHelperResolution,
|
||||
@@ -32920,6 +33158,76 @@ let target_exposure = csi_ready ? dynamic_exposure : 0.0;
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn numeric_vm_compiles_static_runtime_helper_arguments() {
|
||||
let strategy = PlatformExprStrategy::new(PlatformExprStrategyConfig::microcap_rotation());
|
||||
let plan = strategy.expression_eval_plan(
|
||||
"rolling_mean(\"close\", 5) + rolling_mean_current(\"volume\", 10) + \
|
||||
rolling_max_current(\"close\", 20) + rolling_return_stddev_current(\"close\", 30) + \
|
||||
vma(60) + rolling_sum(\"amount\", 5) + rolling_min(\"close\", 10) + \
|
||||
rolling_max(\"close\", 10) + rolling_stddev(\"close\", 20) + \
|
||||
rolling_zscore(\"close\", 20) + pct_change(\"close\", 5) + \
|
||||
factor_value(\"quality_score\", 1)",
|
||||
);
|
||||
let vm = plan.numeric_vm.as_ref().expect("numeric VM plan");
|
||||
let bindings = vm
|
||||
.helper_bindings
|
||||
.iter()
|
||||
.flatten()
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
assert_eq!(bindings.len(), 12);
|
||||
assert!(
|
||||
bindings
|
||||
.iter()
|
||||
.all(|binding| binding.compiled_args.is_some()),
|
||||
"all static numeric helper arguments must be parsed once at plan compilation"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "manual release-mode runtime helper binding benchmark"]
|
||||
fn benchmark_compiled_runtime_helper_arguments() {
|
||||
let args = vec!["\"close\"".to_string(), "20".to_string()];
|
||||
let compiled = PlatformExprStrategy::compile_runtime_helper_args("rolling_mean", &args)
|
||||
.expect("compiled helper arguments");
|
||||
let iterations = 2_000_000usize;
|
||||
|
||||
let generic_started = std::time::Instant::now();
|
||||
let mut generic_checksum = 0usize;
|
||||
for _ in 0..iterations {
|
||||
let field = PlatformExprStrategy::parse_string_or_identifier(std::hint::black_box(
|
||||
&args[0],
|
||||
))
|
||||
.expect("field");
|
||||
let lookback = PlatformExprStrategy::parse_positive_usize(std::hint::black_box(
|
||||
&args[1],
|
||||
))
|
||||
.expect("lookback");
|
||||
generic_checksum = generic_checksum.wrapping_add(field.len() + lookback);
|
||||
}
|
||||
let generic_seconds = generic_started.elapsed().as_secs_f64();
|
||||
|
||||
let compiled_started = std::time::Instant::now();
|
||||
let mut compiled_checksum = 0usize;
|
||||
for _ in 0..iterations {
|
||||
let CompiledRuntimeHelperArgs::RollingMean {
|
||||
field, lookback, ..
|
||||
} = std::hint::black_box(&compiled)
|
||||
else {
|
||||
panic!("unexpected compiled helper binding");
|
||||
};
|
||||
compiled_checksum = compiled_checksum.wrapping_add(field.len() + *lookback);
|
||||
}
|
||||
let compiled_seconds = compiled_started.elapsed().as_secs_f64();
|
||||
|
||||
assert_eq!(generic_checksum, compiled_checksum);
|
||||
println!(
|
||||
"runtime_helper_binding_benchmark iterations={iterations} generic_seconds={generic_seconds:.6} compiled_seconds={compiled_seconds:.6} speedup={:.3} checksum={generic_checksum}",
|
||||
generic_seconds / compiled_seconds.max(f64::EPSILON),
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn numeric_vm_reuses_rolling_helper_program_across_dates() {
|
||||
let dates = [d(2025, 2, 3), d(2025, 2, 4)];
|
||||
|
||||
Reference in New Issue
Block a user