use super::{OptimizerChoice,OptimizerChoiceState,SimpleOptimizer};
use hashbrown::HashMap;
use ruda_model::{module::ParamId,tensor::backend::AutodiffBackend};
use crate::record::{AdaptorRecord,AdaptorRecordV1};
macro_rules! wrap_record {
($record:expr,$side:ident;$($rank:ident),+) => {
match $record {AdaptorRecord::V1(record)=>AdaptorRecord::V1(match record {
$(AdaptorRecordV1::$rank(state)=>AdaptorRecordV1::$rank(OptimizerChoiceState::$side(state))),+
})}
};
}
macro_rules! unwrap_record {
($record:expr,$side:ident;$($rank:ident),+) => {
match $record {AdaptorRecord::V1(record)=>match record {
$(AdaptorRecordV1::$rank(OptimizerChoiceState::$side(state))=>Ok(AdaptorRecord::V1(AdaptorRecordV1::$rank(state))),
original@AdaptorRecordV1::$rank(_)=>Err(AdaptorRecord::V1(original))),+
}}
};
}
macro_rules! record_is {
($record:expr,$side:ident;$($rank:ident),+) => {
match $record {AdaptorRecord::V1(record)=>match record {
$(AdaptorRecordV1::$rank(OptimizerChoiceState::$side(_))=>true),+,
_=>false,
}}
};
}
impl<L,R> OptimizerChoice<L,R> {
pub fn wrap_left_record<B>(record:AdaptorRecord<L,B>) -> AdaptorRecord<Self,B>
where B:AutodiffBackend,L:SimpleOptimizer<B::InnerBackend>,R:SimpleOptimizer<B::InnerBackend> {
wrap_record!(record,Left;Rank0,Rank1,Rank2,Rank3,Rank4,Rank5,Rank6,Rank7,Rank8)
}
pub fn wrap_right_record<B>(record:AdaptorRecord<R,B>) -> AdaptorRecord<Self,B>
where B:AutodiffBackend,L:SimpleOptimizer<B::InnerBackend>,R:SimpleOptimizer<B::InnerBackend> {
wrap_record!(record,Right;Rank0,Rank1,Rank2,Rank3,Rank4,Rank5,Rank6,Rank7,Rank8)
}
pub fn try_unwrap_left_record<B>(record:AdaptorRecord<Self,B>) -> Result<AdaptorRecord<L,B>,AdaptorRecord<Self,B>>
where B:AutodiffBackend,L:SimpleOptimizer<B::InnerBackend>,R:SimpleOptimizer<B::InnerBackend> {
unwrap_record!(record,Left;Rank0,Rank1,Rank2,Rank3,Rank4,Rank5,Rank6,Rank7,Rank8)
}
pub fn try_unwrap_right_record<B>(record:AdaptorRecord<Self,B>) -> Result<AdaptorRecord<R,B>,AdaptorRecord<Self,B>>
where B:AutodiffBackend,L:SimpleOptimizer<B::InnerBackend>,R:SimpleOptimizer<B::InnerBackend> {
unwrap_record!(record,Right;Rank0,Rank1,Rank2,Rank3,Rank4,Rank5,Rank6,Rank7,Rank8)
}
pub fn wrap_left_records<B>(records:HashMap<ParamId,AdaptorRecord<L,B>>) -> HashMap<ParamId,AdaptorRecord<Self,B>>
where B:AutodiffBackend,L:SimpleOptimizer<B::InnerBackend>,R:SimpleOptimizer<B::InnerBackend> {
records.into_iter().map(|(id,record)|(id,Self::wrap_left_record::<B>(record))).collect()
}
pub fn wrap_right_records<B>(records:HashMap<ParamId,AdaptorRecord<R,B>>) -> HashMap<ParamId,AdaptorRecord<Self,B>>
where B:AutodiffBackend,L:SimpleOptimizer<B::InnerBackend>,R:SimpleOptimizer<B::InnerBackend> {
records.into_iter().map(|(id,record)|(id,Self::wrap_right_record::<B>(record))).collect()
}
pub fn try_unwrap_left_records<B>(records:HashMap<ParamId,AdaptorRecord<Self,B>>)
-> Result<HashMap<ParamId,AdaptorRecord<L,B>>,HashMap<ParamId,AdaptorRecord<Self,B>>>
where B:AutodiffBackend,L:SimpleOptimizer<B::InnerBackend>,R:SimpleOptimizer<B::InnerBackend> {
if records.values().any(|record|!record_is!(record,Left;Rank0,Rank1,Rank2,Rank3,Rank4,Rank5,Rank6,Rank7,Rank8)) {return Err(records);}
Ok(records.into_iter().map(|(id,record)| {
let record=Self::try_unwrap_left_record::<B>(record).unwrap_or_else(|_|unreachable!("validated original left algorithm branch"));(id,record)
}).collect())
}
pub fn try_unwrap_right_records<B>(records:HashMap<ParamId,AdaptorRecord<Self,B>>)
-> Result<HashMap<ParamId,AdaptorRecord<R,B>>,HashMap<ParamId,AdaptorRecord<Self,B>>>
where B:AutodiffBackend,L:SimpleOptimizer<B::InnerBackend>,R:SimpleOptimizer<B::InnerBackend> {
if records.values().any(|record|!record_is!(record,Right;Rank0,Rank1,Rank2,Rank3,Rank4,Rank5,Rank6,Rank7,Rank8)) {return Err(records);}
Ok(records.into_iter().map(|(id,record)| {
let record=Self::try_unwrap_right_record::<B>(record).unwrap_or_else(|_|unreachable!("validated original right algorithm branch"));(id,record)
}).collect())
}
}