Skip to main content

ruprim/warp/
mod.rs

1//! Collectives over logical warps within a native device subgroup.
2//!
3//! All lanes of a logical warp participate in each call. `width` is positive,
4//! does not exceed the native subgroup width, and groups must not straddle a
5//! native subgroup. No NVIDIA warp width is assumed by these algorithms.
6
7use 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/// Compute both scan forms with one scan and return the unseeded aggregate
27/// to every participating logical lane.
28#[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/// Inclusive scan in lane order. Values beyond `valid_lanes` are unspecified.
45/// `valid_lanes` is uniform within a logical warp and lies in `1..=width`.
46#[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/// Exclusive scan seeded by `initial`, with the same participation contract.
69#[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/// Unseeded exclusive scan; the first logical lane's output is unspecified.
88#[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/// Ordered reduction, broadcast to all lanes of the logical warp.
98#[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/// Broadcast from a logical lane, not a native subgroup lane.
111#[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/// Segmented reduction whose result is valid at each segment's head lane.
118/// The first lane implicitly starts a segment.
119#[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/// Segmented reduction whose result is valid at each segment's head lane.
149/// A true `tail` marks the last lane of a segment; the final lane is implicit.
150#[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}