use super::*;
use std::collections::HashSet;
pub(super) fn visit_selection<
B: AutodiffBackend,
M: AutodiffModule<B>,
V: ModuleVisitor<B>,
>(model: &M, visitor: &mut V, selection: Option<&[ParamId]>) -> bool {
let Some(selection) = selection else {
model.visit(visitor);
return true;
};
let mut selected = SelectionVisitor {
inner: visitor,
ids: selection.iter().copied().collect(),
seen: HashSet::new(),
};
model.visit(&mut selected);
selected.ids.len() == selection.len() && selected.ids == selected.seen
}
struct SelectionVisitor<'a, V> {
inner: &'a mut V,
ids: HashSet<ParamId>,
seen: HashSet<ParamId>,
}
impl<B: AutodiffBackend, V: ModuleVisitor<B>> ModuleVisitor<B> for SelectionVisitor<'_, V> {
fn enter_module(&mut self, name: &str, container_type: &str) {
self.inner.enter_module(name, container_type);
}
fn exit_module(&mut self, name: &str, container_type: &str) {
self.inner.exit_module(name, container_type);
}
fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B, D>>) {
if self.ids.contains(¶m.id) {
self.seen.insert(param.id);
self.inner.visit_float(param);
}
}
fn visit_int<const D: usize>(&mut self, param: &Param<Tensor<B, D, Int>>) {
if self.ids.contains(¶m.id) {
self.seen.insert(param.id);
self.inner.visit_int(param);
}
}
fn visit_bool<const D: usize>(&mut self, param: &Param<Tensor<B, D, Bool>>) {
if self.ids.contains(¶m.id) {
self.seen.insert(param.id);
self.inner.visit_bool(param);
}
}
}
pub(super) fn map_selection<
B: AutodiffBackend,
M: AutodiffModule<B>,
V: ModuleMapper<B>,
>(model: M, mapper: &mut V, selection: Option<&[ParamId]>) -> M {
let Some(selection) = selection else {
return model.map(mapper);
};
model.map(&mut SelectionMapper {
inner: mapper,
ids: selection.iter().copied().collect(),
})
}
struct SelectionMapper<'a, V> {
inner: &'a mut V,
ids: HashSet<ParamId>,
}
impl<B: AutodiffBackend, V: ModuleMapper<B>> ModuleMapper<B> for SelectionMapper<'_, V> {
fn enter_module(&mut self, name: &str, container_type: &str) {
self.inner.enter_module(name, container_type);
}
fn exit_module(&mut self, name: &str, container_type: &str) {
self.inner.exit_module(name, container_type);
}
fn map_float<const D: usize>(&mut self, param: Param<Tensor<B, D>>) -> Param<Tensor<B, D>> {
if self.ids.contains(¶m.id) {
self.inner.map_float(param)
} else {
param
}
}
fn map_int<const D: usize>(&mut self, param: Param<Tensor<B, D, Int>>) -> Param<Tensor<B, D, Int>> {
if self.ids.contains(¶m.id) {
self.inner.map_int(param)
} else {
param
}
}
fn map_bool<const D: usize>(
&mut self,
param: Param<Tensor<B, D, Bool>>,
) -> Param<Tensor<B, D, Bool>> {
if self.ids.contains(¶m.id) {
self.inner.map_bool(param)
} else {
param
}
}
}
pub(super) fn selected_device_matches<B: AutodiffBackend, M: AutodiffModule<B>>(
model: &M,
device: &B::Device,
selection: Option<&[ParamId]>,
) -> bool {
let Some(selection) = selection else {
return model.devices().iter().all(|actual| actual == device);
};
let mut visitor = SelectedDevice::<B> {
device,
matches: true,
};
let known = visit_selection::<B, _, _>(model, &mut visitor, Some(selection));
known && visitor.matches
}
struct SelectedDevice<'a, B: AutodiffBackend> {
device: &'a B::Device,
matches: bool,
}
impl<B: AutodiffBackend> ModuleVisitor<B> for SelectedDevice<'_, B> {
fn visit_float<const D: usize>(&mut self, param: &Param<Tensor<B, D>>) {
self.matches &= param.val().device() == *self.device;
}
fn visit_int<const D: usize>(&mut self, param: &Param<Tensor<B, D, Int>>) {
self.matches &= param.val().device() == *self.device;
}
fn visit_bool<const D: usize>(&mut self, param: &Param<Tensor<B, D, Bool>>) {
self.matches &= param.val().device() == *self.device;
}
}
#[derive(Debug)]
pub struct SelectedDataParallel<
B: AutodiffBackend,
C: DataParallelCommunicator<B::InnerBackend> = RankCommunicator<
TensorDevice<<B as AutodiffBackend>::InnerBackend>,
>,
> {
inner: DataParallel<B, C>,
parameters: Vec<ParamId>,
}
#[derive(Serialize, Deserialize, PartialEq, Eq)]
struct ReductionMode {
normalize: bool,
fp32: bool,
}
impl<B: AutodiffBackend, C: DataParallelCommunicator<B::InnerBackend>> SelectedDataParallel<B, C> {
pub fn initialize<M: AutodiffModule<B>>(
communicator: C,
model: M,
root: u32,
parameters: &[ParamId],
) -> Result<(Self, M), DataParallelError> {
let (inner, model) =
DataParallel::initialize_inner(communicator, model, root, false, Some(parameters))?;
Ok((
Self {
inner,
parameters: parameters.to_vec(),
},
model,
))
}
pub fn initialize_with_buffers<M: AutodiffModule<B>>(
communicator: C,
model: M,
root: u32,
parameters: &[ParamId],
) -> Result<(Self, M), DataParallelError> {
let (inner, model) =
DataParallel::initialize_inner(communicator, model, root, true, Some(parameters))?;
Ok((Self { inner, parameters: parameters.to_vec() }, model))
}
pub fn rank(&self) -> u32 {
self.inner.rank()
}
pub fn world_size(&self) -> u32 {
self.inner.world_size()
}
pub fn parameters(&self) -> &[ParamId] {
&self.parameters
}
pub(super) fn into_parts(self) -> (DataParallel<B, C>, Vec<ParamId>) {
(self.inner, self.parameters)
}
pub fn reduce<M: AutodiffModule<B>>(
&self,
model: &M,
gradients: GradientsParams,
local_weight: u64,
policy: MissingGradientPolicy,
) -> Result<DataParallelGradients, DataParallelError> {
self.reduce_selected(model, gradients, local_weight, policy, false, true)
}
pub fn reduce_fp32<M: AutodiffModule<B>>(
&self,
model: &M,
gradients: GradientsParams,
local_weight: u64,
policy: MissingGradientPolicy,
) -> Result<DataParallelGradients, DataParallelError> {
self.reduce_selected(model, gradients, local_weight, policy, true, true)
}
pub fn sum<M: AutodiffModule<B>>(
&self,
model: &M,
gradients: GradientsParams,
local_active: bool,
policy: MissingGradientPolicy,
) -> Result<DataParallelGradients, DataParallelError> {
self.reduce_selected(model, gradients, u64::from(local_active), policy, false, false)
}
pub fn sum_fp32<M: AutodiffModule<B>>(
&self,
model: &M,
gradients: GradientsParams,
local_active: bool,
policy: MissingGradientPolicy,
) -> Result<DataParallelGradients, DataParallelError> {
self.reduce_selected(model, gradients, u64::from(local_active), policy, true, false)
}
fn reduce_selected<M: AutodiffModule<B>>(
&self,
model: &M,
gradients: GradientsParams,
local_weight: u64,
policy: MissingGradientPolicy,
fp32: bool,
normalize: bool,
) -> Result<DataParallelGradients, DataParallelError> {
reduce_selected(&self.inner, &self.parameters, model, gradients,
local_weight, policy, fp32, normalize)
}
}
pub(super) fn reduce_selected<
B: AutodiffBackend,
C: DataParallelCommunicator<B::InnerBackend>,
M: AutodiffModule<B>,
>(
session: &DataParallel<B, C>,
parameters: &[ParamId],
model: &M,
gradients: GradientsParams,
local_weight: u64,
policy: MissingGradientPolicy,
fp32: bool,
normalize: bool,
) -> Result<DataParallelGradients, DataParallelError> {
agree_mode(session, fp32, normalize)?;
let (selected, untouched) = gradients.partition::<B::InnerBackend>(parameters);
let reduced = session.reduce_inner(
model, selected, local_weight, policy, fp32, Some(parameters), normalize,
)?;
let gradients = reduced.gradients
.merge_disjoint::<B::InnerBackend>(untouched)
.map_err(|error| contract(error.to_string()))?;
Ok(DataParallelGradients {
gradients,
global_weight: reduced.global_weight,
})
}
pub(super) fn agree_mode<B: AutodiffBackend, C: DataParallelCommunicator<B::InnerBackend>>(
session: &DataParallel<B, C>,
fp32: bool,
normalize: bool,
) -> Result<(), DataParallelError> {
let mode = ReductionMode { normalize, fp32 };
let modes = gather::<B::InnerBackend, C, _>(&session.communicator, &mode)?;
if modes.iter().any(|other| other != &mode) {
return Err(contract("selected replicas disagree on SUM/mean or FP32 reduction"));
}
Ok(())
}