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}