Skip to main content

burn_dispatch/ops/
distributed.rs

1use alloc::vec::Vec;
2
3use burn_backend::{
4    DeviceId,
5    distributed::{
6        CollectiveTensor, DistributedConfig, DistributedOps, DistributedParams, ReduceOperation,
7        TensorRef,
8    },
9    tensor::FloatTensor,
10};
11
12use crate::{Dispatch, DispatchDevice};
13
14macro_rules! dispatch_distributed_devices_arms {
15    (
16        $device:expr,
17        $devices:expr,
18        |$inner_devices:ident| $body:expr;
19        $([$Backend:ident, $cfg:meta]),*
20    ) => {
21        match $device {
22            // Autodiff arm first
23            #[cfg(feature = "autodiff")]
24            $crate::DispatchDevice::Autodiff(inner) => {
25                let inner_devices = $devices
26                    .iter()
27                    .map(|d| {
28                         match &d {
29                            #[cfg(feature = "autodiff")]
30                            $crate::DispatchDevice::Autodiff(d_inner) => *d_inner.inner.clone(),
31                            _ => unreachable!("All devices are expected to be of the same variant."),
32                        }
33                    })
34                    .collect::<Vec<_>>();
35                let inner_devices = inner_devices.as_slice();
36                // Recursively dispatch on inner
37                dispatch_distributed_devices_arms!(
38                    @autodiff
39                    &**inner,
40                    &*inner_devices,
41                    |$inner_devices| $body;
42                    $([$Backend, $cfg]),*
43                )
44            },
45            $(
46                #[cfg($cfg)]
47                $crate::DispatchDevice::$Backend(_) => {
48                    type B = $crate::backends::$Backend;
49                    let $inner_devices = $devices
50                        .iter()
51                        .map(|d| {
52                            let DispatchDevice::$Backend(dev) = d else {
53                                unreachable!("All devices are expected to be of the same variant.")
54                            };
55                            dev.clone()
56                        })
57                        .collect::<Vec<_>>();
58                    $body
59                }
60            )*
61            other => panic!("Distributed operations are not supported for device {other:?}"),
62        }
63    };
64    (
65        @autodiff
66        $device:expr,
67        $devices:expr,
68        |$inner_devices:ident| $body:expr;
69        $([$Backend:ident, $cfg:meta]),*
70    ) => {
71        match $device {
72            $(
73                #[cfg($cfg)]
74                $crate::DispatchDevice::$Backend(_) => {
75                    type B = $crate::backends::Autodiff<$crate::backends::$Backend>;
76                    let $inner_devices = $devices
77                        .iter()
78                        .map(|d| {
79                            let DispatchDevice::$Backend(dev) = d else {
80                                unreachable!("All devices are expected to be of the same variant.")
81                            };
82                            dev.clone()
83                        })
84                        .collect::<Vec<_>>();
85                    $body
86                }
87            )*
88            $crate::DispatchDevice::Autodiff(_) => panic!("Autodiff should not wrap an autodiff device."),
89            other => panic!("Distributed operations are not supported for device {other:?}"),
90        }
91    };
92}
93
94/// Dispatches an operation body based on the provided devices.
95macro_rules! dispatch_distributed_devices {
96    ($device:expr, $devices:expr, |$inner_devices:ident| $body:expr) => {
97        distributed_backend_list!(
98            dispatch_distributed_devices_arms,
99            $device,
100            $devices,
101            |$inner_devices| $body
102        )
103    };
104}
105
106// In builds without a collective-capable backend (Cuda/Remote), the distributed dispatch arms
107// are all cfg'd out, leaving only a diverging fallback — so the captured arguments and trailing
108// expressions are intentionally unused/unreachable.
109#[allow(unused_variables, unreachable_code)]
110impl DistributedOps<Self> for Dispatch {
111    fn start_communication_server(devices: &[DispatchDevice], config: DistributedConfig) {
112        if !devices.is_empty() {
113            let first = &devices[0];
114            dispatch_distributed_devices!(first, devices, |inner_devices| {
115                B::start_communication_server(&inner_devices, config)
116            });
117        }
118    }
119
120    fn close_communication_server(device: &DispatchDevice) {
121        dispatch_device!(@distributed device, |device| {
122            B::close_communication_server(device)
123        })
124    }
125
126    fn register_sync_parameters(
127        device: &DispatchDevice,
128        sharded_param_ids: Vec<DistributedParams>,
129    ) {
130        dispatch_device!(@distributed device, |device| B::register_sync_parameters(
131            device,
132            sharded_param_ids,
133        ))
134    }
135
136    fn submit_sync_collective(device: &DispatchDevice) {
137        dispatch_device!(@distributed device, |device| B::submit_sync_collective(device))
138    }
139
140    fn submit_gradient_sync(_tensor: TensorRef<Self>, _distributed_params: DistributedParams) {
141        unimplemented!()
142    }
143
144    fn all_reduce(
145        tensor: FloatTensor<Self>,
146        op: ReduceOperation,
147        device_ids: Vec<DeviceId>,
148    ) -> CollectiveTensor<Self> {
149        // Safety: we call `assume_resolved` only to wrap it in a new `CollectiveTensor`.
150        // Explicit type: the distributed dispatch only emits arms for collective-capable
151        // backends (Cuda, Remote), so a build with none of them leaves only the diverging
152        // fallback and the match would otherwise infer `!`.
153        let tensor: FloatTensor<Self> = unary_float!(@distributed tensor, float, |tensor| {
154            let collective_tensor = B::all_reduce(tensor, op, device_ids);
155            unsafe { collective_tensor.assume_resolved() }
156        } => Float);
157        CollectiveTensor::new(tensor)
158    }
159
160    fn sync_collective(device: &DispatchDevice) {
161        dispatch_device!(@distributed device, |device| B::sync_collective(device))
162    }
163
164    unsafe fn comm_device(_tensor: &TensorRef<Self>) -> DispatchDevice {
165        unimplemented!()
166    }
167
168    unsafe fn float_from_ref(_tensor: &TensorRef<Self>) -> FloatTensor<Self> {
169        unimplemented!()
170    }
171}