ruccl/in_process/
buffer.rs1use 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}