diff --git a/rust_hft/alpha-harness/domain/src/lib.rs b/rust_hft/alpha-harness/domain/src/lib.rs index 5205731a..92f10a7b 100644 --- a/rust_hft/alpha-harness/domain/src/lib.rs +++ b/rust_hft/alpha-harness/domain/src/lib.rs @@ -2171,6 +2171,7 @@ impl CexBaselineModelV1 { pub enum CexBaselineCartNodeV1 { Leaf { value: f64, + sample_count: usize, }, Split { feature_index: usize, @@ -2333,6 +2334,8 @@ impl CexBaselineArtifactV1 { self.evaluation_policy.validate()?; self.baseline_policy.validate()?; self.evaluation.validate()?; + let (evaluation_protocol, _) = self.evaluation.protocol_binding()?; + validate_baseline_fold_schedule(&self.folds, &evaluation_protocol.walk_forward)?; if self.evaluation.evaluator_version != CEX_BASELINE_WALK_FORWARD_EVALUATOR_VERSION || self.evaluation.protocol_binding()?.1 != self.evaluation_policy.content_sha256 || self.evaluation.formula_config()? != self.baseline_policy.evaluator_config @@ -2397,18 +2400,14 @@ impl CexBaselineArtifactV1 { pub fn validate_binding( &self, mission: &CexResearchMissionArtifactV1, + policy: &CexBaselinePolicyV1, factor_bank: &CexFactorBankRevisionV2, ) -> Result<(), DomainError> { mission.validate()?; + policy.validate()?; factor_bank.validate()?; self.validate()?; - let expected_factor_ids = factor_bank - .entries - .iter() - .map(|entry| entry.factor_id.clone()) - .collect::>(); - let mut expected_factor_ids = expected_factor_ids; - expected_factor_ids.sort(); + let expected_factor_ids = expected_factor_ids(factor_bank); if self.mission_id != mission.semantic_id()? || self.factor_bank_revision_id != factor_bank.revision_id || self.factor_ids != expected_factor_ids @@ -2416,6 +2415,7 @@ impl CexBaselineArtifactV1 { || self.walk_forward_partition != factor_bank.walk_forward_partition || self.evaluation_policy != mission.spec.policies.evaluation || self.target.horizon != mission.spec.instrument.horizon + || self.baseline_policy != *policy || !mission .spec .hypotheses @@ -2424,8 +2424,7 @@ impl CexBaselineArtifactV1 { { return Err(DomainError::InvalidCexBaseline("artifact binding drifted")); } - self.baseline_policy - .validate_binding(&mission.spec.policies.baseline) + policy.validate_binding(&mission.spec.policies.baseline) } fn expected_artifact_id(&self) -> Result { @@ -2454,6 +2453,63 @@ fn fold_ranges_hash(fold: &CexBaselineFoldV1) -> Result { }) } +fn expected_factor_ids(factor_bank: &CexFactorBankRevisionV2) -> Vec { + let mut factor_ids = factor_bank + .entries + .iter() + .map(|entry| entry.factor_id.clone()) + .collect::>(); + factor_ids.sort(); + factor_ids +} + +fn validate_baseline_fold_schedule( + folds: &[CexBaselineFoldV1], + walk_forward: &EvaluationWalkForwardV1, +) -> Result<(), DomainError> { + let fold_step = walk_forward + .validation_rows + .checked_add(walk_forward.embargo_rows) + .ok_or(DomainError::InvalidCexBaseline("fold schedule overflows"))?; + if folds.len() != walk_forward.fold_count { + return Err(DomainError::InvalidCexBaseline( + "fold count is not bound to evaluation protocol", + )); + } + let mut previous_validation_end = None; + for (index, fold) in folds.iter().enumerate() { + let train_end = walk_forward + .initial_train_rows + .checked_add( + index + .checked_mul(fold_step) + .ok_or(DomainError::InvalidCexBaseline("fold schedule overflows"))?, + ) + .ok_or(DomainError::InvalidCexBaseline("fold schedule overflows"))?; + let validation_start = train_end + .checked_add(walk_forward.purge_rows) + .ok_or(DomainError::InvalidCexBaseline("fold schedule overflows"))?; + let validation_end = validation_start + .checked_add(walk_forward.validation_rows) + .ok_or(DomainError::InvalidCexBaseline("fold schedule overflows"))?; + let embargo_end = validation_end + .checked_add(walk_forward.embargo_rows) + .ok_or(DomainError::InvalidCexBaseline("fold schedule overflows"))?; + if fold.train_range != CexBaselineRangeV1::new(0, train_end) + || fold.purge_range != CexBaselineRangeV1::new(train_end, validation_start) + || fold.validation_range != CexBaselineRangeV1::new(validation_start, validation_end) + || fold.embargo_range != CexBaselineRangeV1::new(validation_end, embargo_end) + || previous_validation_end.is_some_and(|end| fold.validation_range.start < end) + { + return Err(DomainError::InvalidCexBaseline( + "fold schedule is not bound to evaluation protocol", + )); + } + previous_validation_end = Some(validation_end); + } + Ok(()) +} + fn artifact_partition_hash(artifact: &CexBaselineArtifactV1) -> Result { let folds = artifact .folds @@ -2509,7 +2565,10 @@ fn validate_cart_node( return Err(DomainError::InvalidCexBaseline("CART exceeds max depth")); } match node { - CexBaselineCartNodeV1::Leaf { value } if value.is_finite() => Ok(()), + CexBaselineCartNodeV1::Leaf { + value, + sample_count, + } if value.is_finite() && *sample_count >= policy.cart_min_leaf => Ok(()), CexBaselineCartNodeV1::Split { feature_index, threshold, @@ -2666,29 +2725,42 @@ impl CexBaselineGateV1 { pub fn validate_binding( &self, + mission: &CexResearchMissionArtifactV1, + policy: &CexBaselinePolicyV1, factor_bank: &CexFactorBankRevisionV2, ridge: Option<&CexBaselineArtifactV1>, cart: Option<&CexBaselineArtifactV1>, ) -> Result<(), DomainError> { + mission.validate()?; + policy.validate()?; self.validate()?; factor_bank.validate()?; - if self.factor_bank_revision_id != factor_bank.revision_id { + let mission_id = mission.semantic_id()?; + policy.validate_binding(&mission.spec.policies.baseline)?; + if self.mission_id != mission_id + || self.policy_id != policy.policy_id + || self.policy_hash != policy.content_hash()? + || self.factor_bank_revision_id != factor_bank.revision_id + { return Err(DomainError::InvalidCexBaseline( - "baseline gate factor bank drifted", + "baseline gate producer binding drifted", )); } match (ridge, cart, &self.ridge_artifact_id, &self.cart_artifact_id) { - (None, None, None, None) if factor_bank.entries.is_empty() => Ok(()), + (None, None, None, None) if factor_bank.entries.is_empty() => { + if Self::empty_factor_bank(mission_id, policy, factor_bank)? != *self { + return Err(DomainError::InvalidCexBaseline( + "empty baseline gate binding drifted", + )); + } + Ok(()) + } (Some(ridge), Some(cart), Some(ridge_id), Some(cart_id)) if ridge.artifact_id == *ridge_id && cart.artifact_id == *cart_id => { - let expected_factor_ids = factor_bank - .entries - .iter() - .map(|entry| entry.factor_id.clone()) - .collect::>() - .into_iter() - .collect::>(); + ridge.validate_binding(mission, policy, factor_bank)?; + cart.validate_binding(mission, policy, factor_bank)?; + let expected_factor_ids = expected_factor_ids(factor_bank); if ridge.factor_ids != expected_factor_ids || cart.factor_ids != expected_factor_ids { return Err(DomainError::InvalidCexBaseline( @@ -6088,7 +6160,7 @@ mod tests { evaluation.score = adjusted_score; evaluation.metrics.adjusted_score = adjusted_score; evaluation.validate().unwrap(); - let ranges = [0, 35, 70].map(|offset| { let t = 200 + offset; (CexBaselineRangeV1::new(0, t), CexBaselineRangeV1::new(t, t + 5), CexBaselineRangeV1::new(t + 5, t + 35), CexBaselineRangeV1::new(t + 35, t + 36)) }); + let ranges = [0, 31, 62].map(|offset| { let t = 200 + offset; (CexBaselineRangeV1::new(0, t), CexBaselineRangeV1::new(t, t + 5), CexBaselineRangeV1::new(t + 5, t + 35), CexBaselineRangeV1::new(t + 35, t + 36)) }); let folds = |model: CexBaselineModelV1| ranges.iter().enumerate().map(|(i, (train, purge, validation, embargo))| CexBaselineFoldV1::new(i + 1, *train, *purge, *validation, *embargo, vec![0.1; 30], model.clone()).unwrap()).collect::>(); let partition_hash = canonical_json_hash(&serde_json::json!({"research_dataset": &bank.research_dataset, "folds": ranges.iter().map(|(train, purge, validation, embargo)| serde_json::json!({"train": train, "purge": purge, "validation": validation, "embargo": embargo})).collect::>() })).unwrap(); let partition = CexResearchContentRefV1 { @@ -6107,7 +6179,7 @@ mod tests { let ridge = CexBaselineArtifactV1::new("mission-1".to_string(), bank.revision_id.clone(), factor_ids, target, bank.research_dataset, partition, bank.evaluation_policy, policy, CexBaselineModelKindV1::Ridge, folds(CexBaselineModelV1::Ridge { intercept: 0.0, means: vec![0.0], scales: vec![1.0], coefficients: vec![0.1] }), evaluation).unwrap(); let mut cart = ridge.clone(); cart.model_kind = CexBaselineModelKindV1::ShallowCart; - for fold in &mut cart.folds { fold.model = CexBaselineModelV1::ShallowCart { root: CexBaselineCartNodeV1::Leaf { value: 0.1 } }; } + for fold in &mut cart.folds { fold.model = CexBaselineModelV1::ShallowCart { root: CexBaselineCartNodeV1::Leaf { value: 0.1, sample_count: 5 } }; } cart.artifact_id = cart.expected_artifact_id().unwrap(); let _gate = CexBaselineGateV1::new(&ridge, &cart).unwrap(); @@ -6134,6 +6206,54 @@ mod tests { gate.failure_codes, [CexBaselineFailureCodeV1::EmptyFactorBank] ); - assert!(gate.validate_binding(&empty, None, None).is_ok()); + let mut control_mission = cex_mission_artifact(Utc::now()); + control_mission.spec.policies.baseline = CexResearchContentRefV1 { + id: policy.policy_id.clone(), + content_sha256: policy.content_hash().unwrap(), + }; + control_mission.validate().unwrap(); + let mission_id = control_mission.semantic_id().unwrap(); + let gate = CexBaselineGateV1::empty_factor_bank(mission_id, &policy, &empty).unwrap(); + assert!(gate + .validate_binding(&control_mission, &policy, &empty, None, None) + .is_ok()); + + let mut mission_drift = control_mission.clone(); + mission_drift.spec.objective.push_str(" drift"); + assert!(gate + .validate_binding(&mission_drift, &policy, &empty, None, None) + .is_err()); + + let mut policy_drift = policy.clone(); + policy_drift.policy_id = "other-baseline-policy".to_string(); + assert!(gate + .validate_binding(&control_mission, &policy_drift, &empty, None, None) + .is_err()); + + let mut leaf_count_drift = cart.clone(); + if let CexBaselineModelV1::ShallowCart { + root: CexBaselineCartNodeV1::Leaf { sample_count, .. }, + } = &mut leaf_count_drift.folds[0].model + { + *sample_count = 4; + } + leaf_count_drift.artifact_id = leaf_count_drift.expected_artifact_id().unwrap(); + assert!(leaf_count_drift.validate().is_err()); + + let mut range_drift = cart; + range_drift.folds[1].validation_range.start += 1; + range_drift.folds[1].validation_range.end += 1; + let hash = fold_ranges_hash(&range_drift.folds[1]).unwrap(); + range_drift.folds[1].fold_id = CexResearchContentRefV1 { + id: format!("cex-baseline-fold-2-{hash}"), + content_sha256: hash, + }; + let partition_hash = artifact_partition_hash(&range_drift).unwrap(); + range_drift.walk_forward_partition = CexResearchContentRefV1 { + id: format!("cex-walk-forward-partition-{partition_hash}"), + content_sha256: partition_hash, + }; + range_drift.artifact_id = range_drift.expected_artifact_id().unwrap(); + assert!(range_drift.validate().is_err()); } } diff --git a/rust_hft/alpha-harness/engine/src/baselines.rs b/rust_hft/alpha-harness/engine/src/baselines.rs index bdabce4f..cdc0f21d 100644 --- a/rust_hft/alpha-harness/engine/src/baselines.rs +++ b/rust_hft/alpha-harness/engine/src/baselines.rs @@ -19,6 +19,128 @@ pub struct CexBaselineRun { pub gate: CexBaselineGateV1, } +pub fn verify_cex_baseline_artifact( + context: &EngineContext<'_>, + factor_bank: &CexFactorBankRevisionV2, + artifact: &CexBaselineArtifactV1, +) -> Result<(), String> { + artifact + .validate() + .map_err(|error| format!("baseline artifact validation failed: {error}"))?; + factor_bank + .validate() + .map_err(|error| format!("factor bank validation failed: {error}"))?; + if artifact.factor_bank_revision_id != factor_bank.revision_id + || artifact.research_dataset != factor_bank.research_dataset + || artifact.evaluation_policy != factor_bank.evaluation_policy + || artifact.walk_forward_partition != factor_bank.walk_forward_partition + { + return Err("baseline artifact Factor Bank binding drifted".to_string()); + } + validate_context_identity( + context, + factor_bank, + &artifact.evaluation_policy, + &artifact.target, + )?; + let (factor_ids, factors) = evaluate_factor_features(context, factor_bank)?; + if artifact.factor_ids != factor_ids { + return Err("baseline artifact factor ordering drifted".to_string()); + } + let features = transpose_factors(&factors, context.rows().len())?; + if artifact.folds.len() != context.folds().len() { + return Err("baseline artifact fold count drifted".to_string()); + } + let labels = labels(context.rows()); + let mut signals = vec![0.0; context.rows().len()]; + let mut ranges = Vec::with_capacity(artifact.folds.len()); + for (fold_index, (fold, context_fold)) in artifact.folds.iter().zip(context.folds()).enumerate() + { + let validation = fold.validation_range.start..fold.validation_range.end; + if fold.train_range.start != context_fold.train.start + || fold.train_range.end != context_fold.train.end + || fold.purge_range.start != context_fold.purge.start + || fold.purge_range.end != context_fold.purge.end + || fold.validation_range.start != context_fold.validation.start + || fold.validation_range.end != context_fold.validation.end + || fold.embargo_range.start != context_fold.embargo.start + || fold.embargo_range.end != context_fold.embargo.end + || validation.end > features.len() + || fold.predictions.len() != validation.len() + { + return Err(format!( + "baseline fold {} validation range drifted", + fold_index + 1 + )); + } + let predictions = match &fold.model { + CexBaselineModelV1::Ridge { .. } => { + let fit = fit_ridge( + &features, + &labels, + context_fold.train.clone(), + artifact.baseline_policy.ridge_l2, + )?; + let refit_model = CexBaselineModelV1::Ridge { + intercept: fit.intercept, + means: fit.means.clone(), + scales: fit.scales.clone(), + coefficients: fit.coefficients.clone(), + }; + if refit_model != fold.model { + return Err(format!( + "baseline fold {} Ridge model drifted", + fold_index + 1 + )); + } + predict_fold_ridge(&fit, &features, &validation)? + } + CexBaselineModelV1::ShallowCart { .. } => { + let tree = fit_cart( + &features, + &labels, + context_fold.train.clone(), + artifact.baseline_policy.cart_max_depth, + artifact.baseline_policy.cart_min_leaf, + &factor_ids, + )?; + let refit_model = CexBaselineModelV1::ShallowCart { + root: cart_node(tree.clone()), + }; + if refit_model != fold.model { + return Err(format!( + "baseline fold {} CART model drifted", + fold_index + 1 + )); + } + predict_fold_cart(&tree, &features, &validation)? + } + }; + if !predictions_equal(&predictions, &fold.predictions) { + return Err(format!( + "baseline fold {} predictions drifted", + fold_index + 1 + )); + } + for (index, prediction) in validation.clone().zip(predictions) { + signals[index] = prediction; + } + ranges.push(validation); + } + let evaluator = FormulaEvaluator::new(artifact.baseline_policy.evaluator_config.clone())?; + let evaluation = evaluator.evaluate_signals( + context.rows(), + &signals, + ranges, + CEX_BASELINE_WALK_FORWARD_EVALUATOR_VERSION, + context.protocol(), + )?; + if evaluation != artifact.evaluation { + return Err("baseline evaluation drifted".to_string()); + } + Ok(()) +} + pub fn evaluate_cex_baselines( context: &EngineContext<'_>, factor_bank: &CexFactorBankRevisionV2, @@ -52,30 +174,7 @@ pub fn evaluate_cex_baselines( gate, }); } - let mut entries = factor_bank.entries.iter().collect::>(); - entries.sort_by(|left, right| left.factor_id.cmp(&right.factor_id)); - let factor_ids = entries - .iter() - .map(|entry| entry.factor_id.clone()) - .collect::>(); - let factors = entries - .iter() - .map(|entry| { - let mut values = evaluate_ast(&entry.canonical_ast, context.rows())?; - if values.iter().any(|value| !value.is_finite()) { - return Err(format!( - "factor {} produced a non-finite value", - entry.factor_id - )); - } - if entry.orientation == CexFactorOrientationV1::Negative { - for value in &mut values { - *value = normalize_zero(-*value); - } - } - Ok(values) - }) - .collect::>, String>>()?; + let (factor_ids, factors) = evaluate_factor_features(context, factor_bank)?; let feature_rows = transpose_factors(&factors, context.rows().len())?; let ridge = fit_artifact( context, @@ -169,6 +268,44 @@ fn validate_fold_range(fold: &WalkForwardFold, row_count: usize) -> Result<(), S Ok(()) } +fn evaluate_factor_features( + context: &EngineContext<'_>, + factor_bank: &CexFactorBankRevisionV2, +) -> Result<(Vec, Vec>), String> { + let mut entries = factor_bank.entries.iter().collect::>(); + entries.sort_by(|left, right| left.factor_id.cmp(&right.factor_id)); + let factor_ids = entries + .iter() + .map(|entry| entry.factor_id.clone()) + .collect::>(); + let factors = evaluate_factor_features_from_entries(context, &entries)?; + Ok((factor_ids, factors)) +} + +fn evaluate_factor_features_from_entries( + context: &EngineContext<'_>, + entries: &[&alpha_domain::CexFactorBankEntryV1], +) -> Result>, String> { + entries + .iter() + .map(|entry| { + let mut values = evaluate_ast(&entry.canonical_ast, context.rows())?; + if values.iter().any(|value| !value.is_finite()) { + return Err(format!( + "factor {} produced a non-finite value", + entry.factor_id + )); + } + if entry.orientation == CexFactorOrientationV1::Negative { + for value in &mut values { + *value = normalize_zero(-*value); + } + } + Ok(values) + }) + .collect::>, String>>() +} + fn transpose_factors(factors: &[Vec], row_count: usize) -> Result>, String> { if factors.is_empty() || factors.iter().any(|factor| factor.len() != row_count) { return Err("Factor Bank feature matrix is invalid".to_string()); @@ -259,7 +396,7 @@ fn fit_artifact( BaselineKind::Ridge => CexBaselineModelKindV1::Ridge, BaselineKind::ShallowCart => CexBaselineModelKindV1::ShallowCart, }; - CexBaselineArtifactV1::new( + let artifact = CexBaselineArtifactV1::new( mission_id.to_string(), factor_bank.revision_id.clone(), factor_ids, @@ -272,7 +409,9 @@ fn fit_artifact( folds, evaluation, ) - .map_err(|error| format!("baseline artifact validation failed: {error}")) + .map_err(|error| format!("baseline artifact validation failed: {error}"))?; + verify_cex_baseline_artifact(context, factor_bank, &artifact)?; + Ok(artifact) } fn labels(rows: &[crate::evaluation::ResearchRow]) -> Vec { @@ -308,9 +447,23 @@ fn predict_fold_cart( .collect() } +fn predictions_equal(left: &[f64], right: &[f64]) -> bool { + left.len() == right.len() + && left + .iter() + .zip(right) + .all(|(left, right)| left.to_bits() == right.to_bits()) +} + fn cart_node(node: CartNode) -> CexBaselineCartNodeV1 { match node { - CartNode::Leaf { value } => CexBaselineCartNodeV1::Leaf { value }, + CartNode::Leaf { + value, + sample_count, + } => CexBaselineCartNodeV1::Leaf { + value, + sample_count, + }, CartNode::Split { feature_index, threshold, @@ -337,6 +490,7 @@ pub(crate) struct RidgeFit { pub(crate) enum CartNode { Leaf { value: f64, + sample_count: usize, }, Split { feature_index: usize, @@ -519,6 +673,7 @@ fn fit_cart_node( if depth >= max_depth || indices.len() < min_leaf.saturating_mul(2) { return CartNode::Leaf { value: normalize_zero(value), + sample_count: indices.len(), }; } @@ -575,6 +730,7 @@ fn fit_cart_node( let Some((_, feature_index, threshold)) = best else { return CartNode::Leaf { value: normalize_zero(value), + sample_count: indices.len(), }; }; let (left, right): (Vec<_>, Vec<_>) = indices @@ -629,7 +785,7 @@ pub(crate) fn predict_cart(model: &CartNode, features: &[f64]) -> Result *value, + CartNode::Leaf { value, .. } => *value, CartNode::Split { feature_index, threshold, diff --git a/rust_hft/alpha-harness/engine/src/formula_evaluator.rs b/rust_hft/alpha-harness/engine/src/formula_evaluator.rs index 954782d0..a4cfbe28 100644 --- a/rust_hft/alpha-harness/engine/src/formula_evaluator.rs +++ b/rust_hft/alpha-harness/engine/src/formula_evaluator.rs @@ -5,7 +5,8 @@ use crate::{ }; use alpha_domain::{ CandidateArtifact, EvaluationCostsV1, EvaluationProtocolV1, - ONNX_WALK_FORWARD_EVALUATOR_VERSION, SEALED_HOLDOUT_EVALUATOR_VERSION, + CEX_BASELINE_WALK_FORWARD_EVALUATOR_VERSION, ONNX_WALK_FORWARD_EVALUATOR_VERSION, + SEALED_HOLDOUT_EVALUATOR_VERSION, }; pub use alpha_domain::{ FormulaEvaluatorConfig, MultipleTestingAdjustment, WALK_FORWARD_EVALUATOR_VERSION, @@ -124,7 +125,9 @@ impl FormulaEvaluator { let mut failures = Vec::new(); let require_icir = matches!( evaluator_version, - WALK_FORWARD_EVALUATOR_VERSION | ONNX_WALK_FORWARD_EVALUATOR_VERSION + WALK_FORWARD_EVALUATOR_VERSION + | CEX_BASELINE_WALK_FORWARD_EVALUATOR_VERSION + | ONNX_WALK_FORWARD_EVALUATOR_VERSION ); if predictive .time_series_ic @@ -1125,6 +1128,46 @@ mod tests { result.validate().unwrap(); } + #[test] + fn cex_baseline_requires_icir_and_rank_icir_when_they_are_missing() { + let input = prepare_dataset(rows(0.0), &protocol(0.0, 1)).unwrap(); + let evaluator = FormulaEvaluator::new(FormulaEvaluatorConfig::default()).unwrap(); + let context = input.engine_context(); + let signals = context + .rows() + .iter() + .map(|row| row.signal) + .collect::>(); + let ranges = context + .folds() + .iter() + .map(|fold| fold.validation.clone()) + .collect::>(); + + let result = evaluator + .evaluate_signals( + context.rows(), + &signals, + ranges, + CEX_BASELINE_WALK_FORWARD_EVALUATOR_VERSION, + context.protocol(), + ) + .unwrap(); + + assert!(!result.passed); + assert_eq!(result.metrics.predictive.time_series_icir, None); + assert_eq!(result.metrics.predictive.time_series_rank_icir, None); + assert!(result + .failure_reasons + .iter() + .any(|reason| reason.starts_with("time-series ICIR"))); + assert!(result + .failure_reasons + .iter() + .any(|reason| reason.starts_with("time-series RankICIR"))); + result.validate().unwrap(); + } + #[test] fn walk_forward_rejects_a_formula_that_live_cannot_construct() { let evaluator = FormulaEvaluator::new(FormulaEvaluatorConfig::default()).unwrap();