Skip to main content

cubecl_cpp/shared/
warp.rs

1use std::fmt::Display;
2
3use crate::shared::{Builtin, Component, FmtLeft};
4
5use super::{Dialect, IndexedValue, Item, Value};
6
7#[derive(Clone, Debug)]
8pub enum WarpInstruction<D: Dialect> {
9    ReduceSum {
10        input: Value<D>,
11        out: Value<D>,
12    },
13    InclusiveSum {
14        input: Value<D>,
15        out: Value<D>,
16    },
17    ExclusiveSum {
18        input: Value<D>,
19        out: Value<D>,
20    },
21    ReduceProd {
22        input: Value<D>,
23        out: Value<D>,
24    },
25    InclusiveProd {
26        input: Value<D>,
27        out: Value<D>,
28    },
29    ExclusiveProd {
30        input: Value<D>,
31        out: Value<D>,
32    },
33    ReduceMax {
34        input: Value<D>,
35        out: Value<D>,
36    },
37    ReduceMin {
38        input: Value<D>,
39        out: Value<D>,
40    },
41    ElectFallback {
42        out: Value<D>,
43    },
44    Elect {
45        out: Value<D>,
46    },
47    All {
48        input: Value<D>,
49        out: Value<D>,
50    },
51    Any {
52        input: Value<D>,
53        out: Value<D>,
54    },
55    Ballot {
56        input: Value<D>,
57        out: Value<D>,
58    },
59    Broadcast {
60        input: Value<D>,
61        id: Value<D>,
62        out: Value<D>,
63    },
64    Shuffle {
65        input: Value<D>,
66        src_lane: Value<D>,
67        out: Value<D>,
68    },
69    ShuffleXor {
70        input: Value<D>,
71        mask: Value<D>,
72        out: Value<D>,
73    },
74    ShuffleUp {
75        input: Value<D>,
76        delta: Value<D>,
77        out: Value<D>,
78    },
79    ShuffleDown {
80        input: Value<D>,
81        delta: Value<D>,
82        out: Value<D>,
83    },
84}
85
86impl<D: Dialect> Display for WarpInstruction<D> {
87    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
88        match self {
89            WarpInstruction::ReduceSum { input, out } => D::warp_reduce_sum(f, input, out),
90            WarpInstruction::ReduceProd { input, out } => D::warp_reduce_prod(f, input, out),
91            WarpInstruction::ReduceMax { input, out } => D::warp_reduce_max(f, input, out),
92            WarpInstruction::ReduceMin { input, out } => D::warp_reduce_min(f, input, out),
93            WarpInstruction::All { input, out } => D::warp_reduce_all(f, input, out),
94            WarpInstruction::Any { input, out } => D::warp_reduce_any(f, input, out),
95
96            WarpInstruction::InclusiveSum { input, out } => {
97                D::warp_reduce_sum_inclusive(f, input, out)
98            }
99            WarpInstruction::InclusiveProd { input, out } => {
100                D::warp_reduce_prod_inclusive(f, input, out)
101            }
102            WarpInstruction::ExclusiveSum { input, out } => {
103                D::warp_reduce_sum_exclusive(f, input, out)
104            }
105            WarpInstruction::ExclusiveProd { input, out } => {
106                D::warp_reduce_prod_exclusive(f, input, out)
107            }
108            WarpInstruction::Ballot { input, out } => {
109                assert_eq!(
110                    input.item().vectorization(),
111                    1,
112                    "Ballot can't support vectorized input"
113                );
114                let out_fmt = out.fmt_left();
115                write!(
116                    f,
117                    "
118{out_fmt} = {{ "
119                )?;
120                D::compile_warp_ballot(f, input, out.item().elem())?;
121                writeln!(f, ", 0, 0, 0 }};")
122            }
123            WarpInstruction::Broadcast { input, id, out } => reduce_broadcast(f, input, out, id),
124            WarpInstruction::Shuffle {
125                input,
126                src_lane,
127                out,
128            } => {
129                let out_fmt = out.fmt_left();
130                write!(f, "{out_fmt} = {{ ")?;
131                for i in 0..input.item().vectorization() {
132                    let comma = if i > 0 { ", " } else { "" };
133                    write!(f, "{comma}")?;
134                    D::compile_warp_shuffle(
135                        f,
136                        &format!("{}", input.index(i)),
137                        input.item().elem(),
138                        &format!("{src_lane}"),
139                    )?;
140                }
141                writeln!(f, " }};")
142            }
143            WarpInstruction::ShuffleXor { input, mask, out } => {
144                let out_fmt = out.fmt_left();
145                write!(f, "{out_fmt} = {{ ")?;
146                for i in 0..input.item().vectorization() {
147                    let comma = if i > 0 { ", " } else { "" };
148                    write!(f, "{comma}")?;
149                    D::compile_warp_shuffle_xor(
150                        f,
151                        &format!("{}", input.index(i)),
152                        input.item().elem(),
153                        &format!("{mask}"),
154                    )?;
155                }
156                writeln!(f, " }};")
157            }
158            WarpInstruction::ShuffleUp { input, delta, out } => {
159                let out_fmt = out.fmt_left();
160                write!(f, "{out_fmt} = {{ ")?;
161                for i in 0..input.item().vectorization() {
162                    let comma = if i > 0 { ", " } else { "" };
163                    write!(f, "{comma}")?;
164                    D::compile_warp_shuffle_up(
165                        f,
166                        &format!("{}", input.index(i)),
167                        input.item().elem(),
168                        &format!("{delta}"),
169                    )?;
170                }
171                writeln!(f, " }};")
172            }
173            WarpInstruction::ShuffleDown { input, delta, out } => {
174                let out_fmt = out.fmt_left();
175                write!(f, "{out_fmt} = {{ ")?;
176                for i in 0..input.item().vectorization() {
177                    let comma = if i > 0 { ", " } else { "" };
178                    write!(f, "{comma}")?;
179                    D::compile_warp_shuffle_down(
180                        f,
181                        &format!("{}", input.index(i)),
182                        input.item().elem(),
183                        &format!("{delta}"),
184                    )?;
185                }
186                writeln!(f, " }};")
187            }
188            WarpInstruction::ElectFallback { out } => {
189                let out = out.fmt_left();
190                write!(
191                    f,
192                    "
193unsigned int mask = __activemask();
194unsigned int leader = __ffs(mask) - 1;
195{out} = threadIdx.x % warpSize == leader;
196            "
197                )
198            }
199            WarpInstruction::Elect { out } => {
200                let out = out.fmt_left();
201                D::compile_warp_elect(f, &out)
202            }
203        }
204    }
205}
206
207pub(crate) fn reduce_operator<D: Dialect>(
208    f: &mut core::fmt::Formatter<'_>,
209    input: &Value<D>,
210    out: &Value<D>,
211    op: &str,
212) -> core::fmt::Result {
213    let in_optimized = input.optimized();
214    let acc_item = in_optimized.item();
215
216    reduce_with_loop(f, input, out, acc_item, |f, acc, index| {
217        let acc_indexed = maybe_index(acc, index);
218        write!(f, "{acc_indexed} {op} ")?;
219        D::compile_warp_shuffle_xor(f, &acc_indexed, acc.item().elem(), "offset")?;
220        writeln!(f, ";")
221    })
222}
223
224pub(crate) fn reduce_comparison<
225    D: Dialect,
226    I: Fn(&mut core::fmt::Formatter<'_>, Item<D>) -> std::fmt::Result,
227>(
228    f: &mut core::fmt::Formatter<'_>,
229    input: &Value<D>,
230    out: &Value<D>,
231    instruction: I,
232) -> core::fmt::Result {
233    let in_optimized = input.optimized();
234    let acc_item = in_optimized.item();
235    reduce_with_loop(f, input, out, acc_item, |f, acc, index| {
236        let acc_indexed = maybe_index(acc, index);
237        let acc_elem = acc_item.elem();
238        write!(f, "        {acc_indexed} = ")?;
239        instruction(f, in_optimized.item())?;
240        write!(f, "({acc_indexed}, ")?;
241        D::compile_warp_shuffle_xor(f, &acc_indexed, acc_elem, "offset")?;
242        writeln!(f, ");")
243    })
244}
245
246pub(crate) fn reduce_inclusive<D: Dialect>(
247    f: &mut core::fmt::Formatter<'_>,
248    input: &Value<D>,
249    out: &Value<D>,
250    op: &str,
251) -> core::fmt::Result {
252    let in_optimized = input.optimized();
253    let acc_item = in_optimized.item();
254
255    reduce_with_loop(f, input, out, acc_item, |f, acc, index| {
256        let acc_indexed = maybe_index(acc, index);
257        let tmp = Value::tmp(Item::Scalar(*acc_item.elem()));
258        let tmp_left = tmp.fmt_left();
259        let lane_id = Builtin::<D>::UnitPosPlane;
260        write!(
261            f,
262            "
263{tmp_left} = "
264        )?;
265        D::compile_warp_shuffle_up(f, &acc_indexed, acc_item.elem(), "offset")?;
266        write!(
267            f,
268            ";
269if({lane_id} >= offset) {{
270    {acc_indexed} {op} {tmp};
271}}
272"
273        )
274    })
275}
276
277pub(crate) fn reduce_exclusive<D: Dialect>(
278    f: &mut core::fmt::Formatter<'_>,
279    input: &Value<D>,
280    out: &Value<D>,
281    op: &str,
282    default: &str,
283) -> core::fmt::Result {
284    let in_optimized = input.optimized();
285    let acc_item = in_optimized.item();
286
287    let inclusive = Value::tmp(acc_item);
288    reduce_inclusive(f, input, &inclusive, op)?;
289    let shfl = Value::tmp(acc_item);
290    writeln!(f, "{} = {{", shfl.fmt_left())?;
291    for k in 0..acc_item.vectorization() {
292        let inclusive_indexed = maybe_index(&inclusive, k);
293        let comma = if k > 0 { ", " } else { "" };
294        write!(f, "{comma}")?;
295        D::compile_warp_shuffle_up(f, &inclusive_indexed.to_string(), acc_item.elem(), "1")?;
296    }
297    writeln!(f, "}};")?;
298    let lane_id = Builtin::<D>::UnitPosPlane;
299
300    write!(
301        f,
302        "{} = ({lane_id} == 0) ? {}{{",
303        out.fmt_left(),
304        out.item(),
305    )?;
306    for _ in 0..out.item().vectorization() {
307        write!(f, "{default},")?;
308    }
309    writeln!(f, "}} : {};", cast(&shfl, out.item()))
310}
311
312pub(crate) fn reduce_broadcast<D: Dialect>(
313    f: &mut core::fmt::Formatter<'_>,
314    input: &Value<D>,
315    out: &Value<D>,
316    id: &Value<D>,
317) -> core::fmt::Result {
318    let out_fmt = out.fmt_left();
319    write!(f, "{out_fmt} = {{ ")?;
320    for i in 0..input.item().vectorization() {
321        let comma = if i > 0 { ", " } else { "" };
322        write!(f, "{comma}")?;
323        D::compile_warp_shuffle(
324            f,
325            &format!("{}", input.index(i)),
326            input.item().elem(),
327            &format!("{id}"),
328        )?;
329    }
330    writeln!(f, " }};")
331}
332
333fn reduce_with_loop<
334    D: Dialect,
335    I: Fn(&mut core::fmt::Formatter<'_>, &Value<D>, usize) -> std::fmt::Result,
336>(
337    f: &mut core::fmt::Formatter<'_>,
338    input: &Value<D>,
339    out: &Value<D>,
340    acc_item: Item<D>,
341    instruction: I,
342) -> core::fmt::Result {
343    let acc = Value::tmp(acc_item);
344    let vectorization = acc_item.vectorization();
345
346    writeln!(f, "auto plane_{out} = [&]() -> {} {{", out.item())?;
347    writeln!(f, "    {} {} = {};", acc_item, acc, cast(input, acc_item))?;
348    write!(f, "    for (uint offset = 1; offset < ")?;
349    D::compile_plane_dim_checked(f)?;
350    writeln!(f, "; offset *=2 ) {{")?;
351    for k in 0..vectorization {
352        instruction(f, &acc, k)?;
353    }
354    writeln!(f, "    }};")?;
355    writeln!(f, "    return {};", cast(&acc, out.item()))?;
356    writeln!(f, "}};")?;
357    writeln!(f, "{} = plane_{}();", out.fmt_left(), out)
358}
359
360pub(crate) fn reduce_quantifier<
361    D: Dialect,
362    Q: Fn(&mut core::fmt::Formatter<'_>, &IndexedValue<D>) -> std::fmt::Result,
363>(
364    f: &mut core::fmt::Formatter<'_>,
365    input: &Value<D>,
366    out: &Value<D>,
367    quantifier: Q,
368) -> core::fmt::Result {
369    let out_fmt = out.fmt_left();
370    write!(f, "{out_fmt} = {{ ")?;
371    for i in 0..input.item().vectorization() {
372        let comma = if i > 0 { ", " } else { "" };
373        write!(f, "{comma}")?;
374        quantifier(f, &input.index(i))?;
375    }
376    writeln!(f, "}};")
377}
378
379fn cast<D: Dialect>(input: &Value<D>, target: Item<D>) -> String {
380    if target != input.item() {
381        let addr_space = D::address_space_for_value(input);
382        let qualifier = input.const_qualifier();
383        format!("reinterpret_cast<{addr_space}{target}{qualifier}&>({input})")
384    } else {
385        format!("{input}")
386    }
387}
388
389fn maybe_index<D: Dialect>(val: &Value<D>, k: usize) -> String {
390    if val.item().vectorization() > 1 {
391        format!("{val}.i_{k}")
392    } else {
393        format!("{val}")
394    }
395}