burn_dispatch/ops/
distributed.rs1use 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 #[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 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
94macro_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#[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 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}