#[derive(Debug, Clone)]
pub struct PolicyCutRecord<'a> {
pub cut_id: u64,
pub slot_index: u32,
pub iteration: u32,
pub forward_pass_index: u32,
pub intercept: f64,
pub coefficients: &'a [f64],
pub is_active: bool,
}
#[derive(Debug, Clone)]
pub struct PolicyBasisRecord<'a> {
pub stage_id: u32,
pub iteration: u32,
pub column_status: &'a [u8],
pub row_status: &'a [u8],
pub num_cut_rows: u32,
}
#[derive(Debug, Clone)]
pub struct StageStatesPayload<'a> {
pub stage_id: u32,
pub state_dimension: u32,
pub count: u32,
pub data: &'a [f64],
}
#[derive(Debug)]
pub struct StageCutsPayload<'a> {
pub stage_id: u32,
pub state_dimension: u32,
pub capacity: u32,
pub warm_start_count: u32,
pub cuts: &'a [PolicyCutRecord<'a>],
pub active_cut_indices: &'a [u32],
pub populated_count: u32,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct PolicyCheckpointMetadata {
pub cobre_version: String,
pub created_at: String,
pub completed_iterations: u32,
pub final_lower_bound: f64,
pub best_upper_bound: Option<f64>,
pub state_dimension: u32,
pub num_stages: u32,
pub max_iterations: u32,
pub forward_passes: u32,
pub warm_start_cuts: u32,
#[serde(default)]
pub warm_start_counts: Vec<u32>,
pub rng_seed: u64,
#[serde(default)]
pub total_visited_states: u64,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct OwnedPolicyCutRecord {
pub cut_id: u64,
pub slot_index: u32,
pub iteration: u32,
pub forward_pass_index: u32,
pub intercept: f64,
pub coefficients: Vec<f64>,
pub is_active: bool,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct OwnedPolicyBasisRecord {
pub stage_id: u32,
pub iteration: u32,
pub column_status: Vec<u8>,
pub row_status: Vec<u8>,
pub num_cut_rows: u32,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct StageCutsReadResult {
pub stage_id: u32,
pub state_dimension: u32,
pub capacity: u32,
pub warm_start_count: u32,
pub populated_count: u32,
pub cuts: Vec<OwnedPolicyCutRecord>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct StageStatesReadResult {
pub stage_id: u32,
pub state_dimension: u32,
pub count: u32,
pub data: Vec<f64>,
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct PolicyCheckpoint {
pub metadata: PolicyCheckpointMetadata,
pub stage_cuts: Vec<StageCutsReadResult>,
pub stage_bases: Vec<OwnedPolicyBasisRecord>,
pub stage_states: Vec<StageStatesReadResult>,
}