use super::*;
use crate::data_parallel::selected::reduce_selected;
pub struct SelectedZero1<
B,
M,
O,
C = RankCommunicator<TensorDevice<<B as AutodiffBackend>::InnerBackend>>,
> where
B: AutodiffBackend,
M: AutodiffModule<B>,
O: SimpleOptimizer<B::InnerBackend>,
C: DataParallelCommunicator<B::InnerBackend>,
{
inner: Zero1<B, M, O, C>,
parameters: Vec<ParamId>,
}
pub struct SelectedZero1Step<M> {
pub model: M,
pub remaining_gradients: GradientsParams,
pub global_weight: u64,
}
pub struct SelectedZero1Record<B, M, O>
where
B: AutodiffBackend,
M: AutodiffModule<B>,
O: SimpleOptimizer<B::InnerBackend>,
{
version: u32,
parameters: Vec<u64>,
schema: String,
optimizer: Zero1Record<B, M, 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> SelectedZero1<B, M, O, C>
where
B: AutodiffBackend,
M: AutodiffModule<B>,
O: SimpleOptimizer<B::InnerBackend>,
C: DataParallelCommunicator<B::InnerBackend>,
{
pub fn new(
session: SelectedDataParallel<B, C>,
model: &M,
optimizer: O,
owners: Vec<u32>,
) -> Result<Self, DataParallelError> {
let (session, parameters) = session.into_parts();
let inner = Zero1::from_session(session, model, optimizer, owners, Some(¶meters))?;
Ok(Self { inner, parameters })
}
pub fn parameters(&self) -> &[ParamId] {
&self.parameters
}
pub fn trainable_parameters(&self) -> &[ParamId] {
&self.inner.ids
}
pub fn owners(&self) -> &[u32] {
self.inner.owners()
}
pub fn rank(&self) -> u32 {
self.inner.session.rank()
}
pub fn world_size(&self) -> u32 {
self.inner.session.world_size()
}
pub fn owned_parameter_count(&self) -> usize {
self.inner.owned_parameter_count()
}
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<SelectedZero1Step<M>, DataParallelError> {
self.step_inner(lr, model, gradients, local_weight, policy, false, true)
}
pub fn step_fp32(
&mut self,
lr: LearningRate,
model: M,
gradients: GradientsParams,
local_weight: u64,
policy: MissingGradientPolicy,
) -> Result<SelectedZero1Step<M>, DataParallelError> {
self.step_inner(lr, model, gradients, local_weight, policy, true, true)
}
pub fn step_sum(
&mut self,
lr: LearningRate,
model: M,
gradients: GradientsParams,
local_active: bool,
policy: MissingGradientPolicy,
) -> Result<SelectedZero1Step<M>, DataParallelError> {
self.step_inner(lr, model, gradients, u64::from(local_active), policy, false, false)
}
pub fn step_sum_fp32(
&mut self,
lr: LearningRate,
model: M,
gradients: GradientsParams,
local_active: bool,
policy: MissingGradientPolicy,
) -> Result<SelectedZero1Step<M>, DataParallelError> {
self.step_inner(lr, model, gradients, u64::from(local_active), policy, true, false)
}
fn step_inner(
&mut self,
lr: LearningRate,
model: M,
gradients: GradientsParams,
local_weight: u64,
policy: MissingGradientPolicy,
fp32: bool,
normalize: bool,
) -> Result<SelectedZero1Step<M>, DataParallelError> {
let rates = gather::<B::InnerBackend, C, _>(
&self.inner.session.communicator, &lr.to_bits(),
)?;
if !lr.is_finite() || lr < 0. || rates.iter().any(|&other| other != lr.to_bits()) {
return Err(contract("selected ZeRO-1 learning rates must be finite, nonnegative and identical"));
}
let reduced = reduce_selected(
&self.inner.session, &self.parameters, &model, gradients,
local_weight, policy, fp32, normalize,
)?;
let (selected, remaining_gradients) =
reduced.gradients.partition::<B::InnerBackend>(&self.parameters);
let model = self.inner.update_owned(lr, model, selected, Some(&self.parameters))?;
Ok(SelectedZero1Step { model, remaining_gradients, global_weight: reduced.global_weight })
}
pub fn to_record(&self) -> SelectedZero1Record<B, M, O> {
SelectedZero1Record {
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: SelectedZero1Record<B, M, 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-1 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, M, O> Record<B> for SelectedZero1Record<B, M, O>
where
B: AutodiffBackend,
M: AutodiffModule<B>,
O: SimpleOptimizer<B::InnerBackend>,
{
type Item<S: PrecisionSettings> =
<(u32, Vec<u64>, String, Zero1Record<B, M, O>) as Record<B>>::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: &B::Device) -> Self {
let (version, parameters, schema, optimizer) = Record::<B>::from_item::<S>(item, device);
Self { version, parameters, schema, optimizer }
}
}