use super::*;
pub struct SelectedZero2<
B,
M,
O,
C = RankCommunicator<TensorDevice<<B as AutodiffBackend>::InnerBackend>>,
> where
B: AutodiffBackend,
M: AutodiffModule<B>,
O: ElementwiseShardOptimizer<B::InnerBackend>,
C: ShardedCommunicator<B::InnerBackend>,
{
inner: Zero2<B, M, O, C>,
parameters: Vec<ParamId>,
}
pub struct SelectedZero2Step<M> {
pub model: M,
pub remaining_gradients: GradientsParams,
pub global_weight: u64,
}
pub struct SelectedZero2Record<B, O>
where
B: AutodiffBackend,
O: ElementwiseShardOptimizer<B::InnerBackend>,
{
version: u32,
parameters: Vec<u64>,
schema: String,
optimizer: Zero2Record<B, O>,
}
fn parameter_set(parameters: &[ParamId]) -> Vec<u64> {
let mut ids = parameters.iter().map(ParamId::val).collect::<Vec<_>>();
ids.sort_unstable();
ids
}
impl<B, M, O, C> SelectedZero2<B, M, O, C>
where
B: AutodiffBackend,
M: AutodiffModule<B>,
O: ElementwiseShardOptimizer<B::InnerBackend>,
C: ShardedCommunicator<B::InnerBackend>,
{
pub fn new(
session: SelectedDataParallel<B, C>,
model: &M,
optimizer: O,
) -> Result<Self, DataParallelError> {
let (session, parameters) = session.into_parts();
let inner = Zero2::from_session(session, model, optimizer, Some(¶meters))?;
Ok(Self { inner, parameters })
}
pub fn parameters(&self) -> &[ParamId] {
&self.parameters
}
pub fn trainable_parameters(&self) -> &[ParamId] {
&self.inner.ids
}
pub fn rank(&self) -> u32 {
self.inner.rank()
}
pub fn world_size(&self) -> u32 {
self.inner.session.world_size()
}
pub fn state_parameter_count(&self) -> usize {
self.inner.state_parameter_count()
}
pub fn step(
&mut self,
lr: LearningRate,
model: M,
gradients: GradientsParams,
local_weight: u64,
policy: MissingGradientPolicy,
) -> Result<SelectedZero2Step<M>, DataParallelError> {
self.step_inner(lr, model, gradients, local_weight, policy, true)
}
pub fn step_sum(
&mut self,
lr: LearningRate,
model: M,
gradients: GradientsParams,
local_active: bool,
policy: MissingGradientPolicy,
) -> Result<SelectedZero2Step<M>, DataParallelError> {
self.step_inner(lr, model, gradients, u64::from(local_active), policy, false)
}
fn step_inner(
&mut self,
lr: LearningRate,
model: M,
gradients: GradientsParams,
local_weight: u64,
policy: MissingGradientPolicy,
normalize: bool,
) -> Result<SelectedZero2Step<M>, DataParallelError> {
let (updated, remaining_gradients) = self.inner.step_inner(
lr, model, gradients, local_weight, policy, Some(&self.parameters), normalize,
)?;
Ok(SelectedZero2Step {
model: updated.model,
remaining_gradients,
global_weight: updated.global_weight,
})
}
pub fn to_record(&self) -> SelectedZero2Record<B, O> {
SelectedZero2Record {
version: 1,
parameters: parameter_set(&self.parameters),
schema: serde_json::to_string(&self.inner.session.contract)
.expect("selected parameter contract serialization failed"),
optimizer: self.inner.to_record(),
}
}
pub fn load_record(
&mut self,
record: SelectedZero2Record<B, O>,
) -> Result<(), DataParallelError> {
let schema = serde_json::to_string(&self.inner.session.contract)
.map_err(|error| contract(error.to_string()))?;
let error = if record.version != 1
|| record.parameters != parameter_set(&self.parameters)
|| record.schema != schema
{
Some("selected ZeRO-2 membership or original replica schema differs".to_string())
} else {
None
};
for error in gather::<B::InnerBackend, C, _>(&self.inner.session.communicator, &error)? {
if let Some(error) = error {
return Err(contract(error));
}
}
self.inner.load_record(record.optimizer)
}
}
impl<B, O> Record<B::InnerBackend> for SelectedZero2Record<B, O>
where
B: AutodiffBackend,
O: ElementwiseShardOptimizer<B::InnerBackend>,
{
type Item<S: PrecisionSettings> =
<(u32, Vec<u64>, String, Zero2Record<B, O>) as Record<B::InnerBackend>>::Item<S>;
fn into_item<S: PrecisionSettings>(self) -> Self::Item<S> {
(self.version, self.parameters, self.schema, self.optimizer).into_item::<S>()
}
fn from_item<S: PrecisionSettings>(
item: Self::Item<S>,
device: &ruda_model::tensor::Device<B::InnerBackend>,
) -> Self {
let (version, parameters, schema, optimizer) =
Record::<B::InnerBackend>::from_item::<S>(item, device);
Self { version, parameters, schema, optimizer }
}
}