1use ruda_kernel::dsl as kernel_dsl;
8use ruda_kernel::dsl::prelude::*;
9use crate::collective::{RudaBinaryOp, RudaBinaryOpExpand};
10
11pub mod exchange;
12pub mod io;
13pub mod sort;
14pub mod batched;
15pub mod record;
16
17#[ruda]
18pub fn logical_lane_id(#[comptime] width: u32) -> u32 { UNIT_POS_PLANE % width }
19
20#[ruda]
21pub fn logical_warp_id(#[comptime] width: u32) -> u32 { UNIT_POS_PLANE / width }
22
23#[ruda]
24pub fn logical_warp_base_id(#[comptime] width: u32) -> u32 { UNIT_POS_PLANE / width * width }
25
26#[ruda]
29pub fn scan<T: RudaPrimitive, O: RudaBinaryOp<T>>(
30 value: T, initial: T, inclusive_output: &mut T, exclusive_output: &mut T,
31 op: &O, valid_lanes: u32, #[comptime] width: u32,
32) -> T {
33 let inclusive = inclusive_scan::<T, O>(value, op, valid_lanes, width);
34 let lane = logical_lane_id(width);
35 let base = logical_warp_base_id(width);
36 let previous = plane_shuffle(inclusive, base + select(lane > 0, lane - 1, lane));
37 *inclusive_output = inclusive;
38 let mut exclusive = initial;
39 if lane > 0 && lane < valid_lanes { exclusive = op.combine(initial, previous); }
40 *exclusive_output = exclusive;
41 plane_shuffle(inclusive, base + valid_lanes - 1)
42}
43
44#[ruda]
47pub fn inclusive_scan<T: RudaPrimitive, O: RudaBinaryOp<T>>(
48 value: T,
49 op: &O,
50 valid_lanes: u32,
51 #[comptime] width: u32,
52) -> T {
53 let lane = UNIT_POS_PLANE % width;
54 let base = UNIT_POS_PLANE - lane;
55 let mut result = value;
56 let mut distance = 1u32;
57 while distance < width {
58 let source = base + select(lane >= distance, lane - distance, lane);
59 let left = plane_shuffle(result, source);
60 if lane >= distance && lane < valid_lanes {
61 result = op.combine(left, result);
62 }
63 distance *= 2;
64 }
65 result
66}
67
68#[ruda]
70pub fn exclusive_scan<T: RudaPrimitive, O: RudaBinaryOp<T>>(
71 value: T,
72 initial: T,
73 op: &O,
74 valid_lanes: u32,
75 #[comptime] width: u32,
76) -> T {
77 let inclusive = inclusive_scan::<T, O>(value, op, valid_lanes, width);
78 let lane = UNIT_POS_PLANE % width;
79 let previous = plane_shuffle(inclusive, UNIT_POS_PLANE - select(lane > 0, 1u32, 0u32));
80 let mut result = initial;
81 if lane > 0 && lane < valid_lanes {
82 result = op.combine(initial, previous);
83 }
84 result
85}
86
87#[ruda]
89pub fn exclusive_unseeded<T: RudaPrimitive, O: RudaBinaryOp<T>>(
90 value: T, op: &O, valid_lanes: u32, #[comptime] width: u32,
91) -> T {
92 let inclusive = inclusive_scan::<T, O>(value, op, valid_lanes, width);
93 let lane = UNIT_POS_PLANE % width;
94 plane_shuffle(inclusive, UNIT_POS_PLANE - select(lane > 0, 1u32, 0u32))
95}
96
97#[ruda]
99pub fn reduce<T: RudaPrimitive, O: RudaBinaryOp<T>>(
100 value: T,
101 op: &O,
102 valid_lanes: u32,
103 #[comptime] width: u32,
104) -> T {
105 let inclusive = inclusive_scan::<T, O>(value, op, valid_lanes, width);
106 let base = UNIT_POS_PLANE - UNIT_POS_PLANE % width;
107 plane_shuffle(inclusive, base + valid_lanes - 1)
108}
109
110#[ruda]
112pub fn broadcast<T: RudaPrimitive>(value: T, source: u32, #[comptime] width: u32) -> T {
113 let base = UNIT_POS_PLANE - UNIT_POS_PLANE % width;
114 plane_shuffle(value, base + source)
115}
116
117#[ruda]
120pub fn head_segmented_reduce<T: RudaPrimitive, O: RudaBinaryOp<T>>(
121 value: T,
122 head: bool,
123 op: &O,
124 #[comptime] width: u32,
125) -> T {
126 let lane = UNIT_POS_PLANE % width;
127 let base = UNIT_POS_PLANE - lane;
128 let mut result = value;
129 let next_lane = select(lane + 1 < width, lane + 1, lane);
130 let next_head = plane_shuffle(head, base + next_lane);
131 let mut boundary = next_head || lane + 1 == width;
132 let mut distance = 1u32;
133 while distance < width {
134 let source = base + select(lane + distance < width, lane + distance, lane);
135 let right = plane_shuffle(result, source);
136 let right_boundary = plane_shuffle(boundary, source);
137 if lane + distance < width {
138 if !boundary {
139 result = op.combine(result, right);
140 }
141 boundary = boundary || right_boundary;
142 }
143 distance *= 2;
144 }
145 result
146}
147
148#[ruda]
151pub fn tail_segmented_reduce<T: RudaPrimitive, O: RudaBinaryOp<T>>(
152 value: T,
153 tail: bool,
154 op: &O,
155 #[comptime] width: u32,
156) -> T {
157 let lane = UNIT_POS_PLANE % width;
158 let base = UNIT_POS_PLANE - lane;
159 let mut result = value;
160 let mut boundary = tail || lane + 1 == width;
161 let mut distance = 1u32;
162 while distance < width {
163 let source = base + select(lane + distance < width, lane + distance, lane);
164 let right = plane_shuffle(result, source);
165 let right_boundary = plane_shuffle(boundary, source);
166 if lane + distance < width {
167 if !boundary {
168 result = op.combine(result, right);
169 }
170 boundary = boundary || right_boundary;
171 }
172 distance *= 2;
173 }
174 result
175}