Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
166 changes: 143 additions & 23 deletions rust_hft/alpha-harness/domain/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2171,6 +2171,7 @@ impl CexBaselineModelV1 {
pub enum CexBaselineCartNodeV1 {
Leaf {
value: f64,
sample_count: usize,
},
Split {
feature_index: usize,
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -2397,25 +2400,22 @@ 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::<Vec<_>>();
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
|| self.research_dataset != factor_bank.research_dataset
|| 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
Expand All @@ -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<String, DomainError> {
Expand Down Expand Up @@ -2454,6 +2453,63 @@ fn fold_ranges_hash(fold: &CexBaselineFoldV1) -> Result<String, DomainError> {
})
}

fn expected_factor_ids(factor_bank: &CexFactorBankRevisionV2) -> Vec<String> {
let mut factor_ids = factor_bank
.entries
.iter()
.map(|entry| entry.factor_id.clone())
.collect::<Vec<_>>();
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<String, DomainError> {
let folds = artifact
.folds
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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::<BTreeSet<_>>()
.into_iter()
.collect::<Vec<_>>();
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(
Expand Down Expand Up @@ -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::<Vec<_>>();
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::<Vec<_>>() })).unwrap();
let partition = CexResearchContentRefV1 {
Expand All @@ -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();

Expand All @@ -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());
}
}
Loading
Loading