use super::{SimpleOptimizer,Adam,AdamW,Sgd,AdaGrad,RmsProp,Adan};
use ruda_model::tensor::{Tensor,DType,BroadcastTensorCollective,backend::Backend};
pub trait ElementwiseShardOptimizer<B:Backend>:SimpleOptimizer<B> {
fn validate_element_sharding(&self) -> Result<(),&'static str> {Ok(())}
fn shard_gradient_dtype(&self,storage:DType) -> DType {storage}
fn validate_fully_sharded_execution(&self) -> Result<(),&'static str> {self.validate_element_sharding()}
fn validate_fully_sharded_history(&self,_state:&Self::State<1>) -> Result<(),&'static str> {Ok(())}
fn step_fully_sharded<C:BroadcastTensorCollective<B>>(&self,lr:crate::LearningRate,tensor:Tensor<B,1>,gradient:Tensor<B,1>,
state:Option<Self::State<1>>, _binding:&crate::FullyShardedOptimizerParameter<C>)
-> Result<(Tensor<B,1>,Option<Self::State<1>>),crate::FullyShardedElementwiseError<C::Error>> {
Ok(self.step(lr,tensor,gradient,state))
}
}
impl<B:Backend> ElementwiseShardOptimizer<B> for Adam {}
impl<B:Backend> ElementwiseShardOptimizer<B> for AdamW {}
impl<B:Backend> ElementwiseShardOptimizer<B> for Sgd<B> {}
impl<B:Backend> ElementwiseShardOptimizer<B> for AdaGrad {}
impl<B:Backend> ElementwiseShardOptimizer<B> for RmsProp {}
impl<B:Backend> ElementwiseShardOptimizer<B> for Adan {}