Skip to main content

ruccl/in_process/
buffer.rs

1use super::device::InProcessDevice;
2use std::fmt::{Debug, Formatter};
3use std::marker::PhantomData;
4
5pub struct DistributedBuffer<T, D: InProcessDevice<T>> {
6    pub(super) communicator_id: u64,
7    pub(super) buffers: Vec<D::Buffer>,
8    pub(super) length_per_rank: usize,
9    pub(super) marker: PhantomData<fn() -> (T, D)>,
10}
11
12impl<T, D: InProcessDevice<T>> Clone for DistributedBuffer<T, D> {
13    fn clone(&self) -> Self {
14        Self {
15            communicator_id: self.communicator_id,
16            buffers: self.buffers.clone(),
17            length_per_rank: self.length_per_rank,
18            marker: PhantomData,
19        }
20    }
21}
22
23impl<T, D> Debug for DistributedBuffer<T, D>
24where
25    D: InProcessDevice<T>,
26    D::Buffer: Debug,
27{
28    fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
29        formatter
30            .debug_struct("DistributedBuffer")
31            .field("communicator_id", &self.communicator_id)
32            .field("buffers", &self.buffers)
33            .field("length_per_rank", &self.length_per_rank)
34            .finish()
35    }
36}
37
38pub struct VariableDistributedBuffer<T, D: InProcessDevice<T>> {
39    pub(super) communicator_id: u64,
40    pub(super) buffers: Vec<Option<D::Buffer>>,
41    pub(super) lengths: Vec<usize>,
42    pub(super) marker: PhantomData<fn() -> (T, D)>,
43}
44
45impl<T, D: InProcessDevice<T>> Clone for VariableDistributedBuffer<T, D> {
46    fn clone(&self) -> Self {
47        Self {
48            communicator_id: self.communicator_id,
49            buffers: self.buffers.clone(),
50            lengths: self.lengths.clone(),
51            marker: PhantomData,
52        }
53    }
54}
55
56impl<T, D> Debug for VariableDistributedBuffer<T, D>
57where
58    D: InProcessDevice<T>,
59    D::Buffer: Debug,
60{
61    fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
62        formatter
63            .debug_struct("VariableDistributedBuffer")
64            .field("communicator_id", &self.communicator_id)
65            .field("buffers", &self.buffers)
66            .field("lengths", &self.lengths)
67            .finish()
68    }
69}
70
71pub struct RootedBuffer<T, D: InProcessDevice<T>> {
72    pub(super) communicator_id: u64,
73    pub(super) root: usize,
74    pub(super) buffer: D::Buffer,
75    pub(super) marker: PhantomData<fn() -> (T, D)>,
76}
77
78impl<T, D: InProcessDevice<T>> Clone for RootedBuffer<T, D> {
79    fn clone(&self) -> Self {
80        Self {
81            communicator_id: self.communicator_id,
82            root: self.root,
83            buffer: self.buffer.clone(),
84            marker: PhantomData,
85        }
86    }
87}
88
89impl<T, D> Debug for RootedBuffer<T, D>
90where
91    D: InProcessDevice<T>,
92    D::Buffer: Debug,
93{
94    fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
95        formatter
96            .debug_struct("RootedBuffer")
97            .field("communicator_id", &self.communicator_id)
98            .field("root", &self.root)
99            .field("buffer", &self.buffer)
100            .finish()
101    }
102}
103
104impl<T, D: InProcessDevice<T>> DistributedBuffer<T, D> {
105    pub fn world_size(&self) -> usize {
106        self.buffers.len()
107    }
108
109    pub fn length_per_rank(&self) -> usize {
110        self.length_per_rank
111    }
112
113    pub fn rank_buffer(&self, rank: usize) -> Option<&D::Buffer> {
114        self.buffers.get(rank)
115    }
116}
117
118impl<T, D: InProcessDevice<T>> VariableDistributedBuffer<T, D> {
119    pub fn world_size(&self) -> usize {
120        self.buffers.len()
121    }
122
123    pub fn rank_length(&self, rank: usize) -> Option<usize> {
124        self.lengths.get(rank).copied()
125    }
126
127    pub fn rank_buffer(&self, rank: usize) -> Option<&D::Buffer> {
128        self.buffers.get(rank).and_then(Option::as_ref)
129    }
130}
131
132impl<T, D: InProcessDevice<T>> RootedBuffer<T, D> {
133    pub const fn root(&self) -> usize {
134        self.root
135    }
136
137    pub fn len(&self) -> usize {
138        D::buffer_len(&self.buffer)
139    }
140
141    pub fn is_empty(&self) -> bool {
142        D::buffer_is_empty(&self.buffer)
143    }
144
145    pub const fn buffer(&self) -> &D::Buffer {
146        &self.buffer
147    }
148}