diff --git a/cuda-oxide/crates/mir-lower/src/convert/ops/arithmetic.rs b/cuda-oxide/crates/mir-lower/src/convert/ops/arithmetic.rs index 0c16f6f7e0..c79032061d 100644 --- a/cuda-oxide/crates/mir-lower/src/convert/ops/arithmetic.rs +++ b/cuda-oxide/crates/mir-lower/src/convert/ops/arithmetic.rs @@ -474,6 +474,14 @@ pub(crate) fn convert_shl( /// (trunc) of the shift amount to match, then masks it with `bit_width - 1`. /// That matches Rust's unchecked/release shift behavior and avoids LLVM poison /// for oversized counts. +/// +/// Two representation rules apply on top of the width rules. The shift value +/// is unified to the signless representation first, because every op here is an +/// `IntBinArithOp` and pliron-llvm rejects a represented operand of one. A count +/// that differs in representation at equal width is then a representation +/// change rather than a width change: `trunc` is illegal there (LLVM requires a +/// strictly smaller result) and reusing the count unchanged would leave the +/// mask and the shift disagreeing with it. fn convert_shift( ctx: &mut Context, rewriter: &mut DialectConversionRewriter, @@ -486,7 +494,6 @@ where let (lhs, rhs) = get_binary_operands(op, ctx)?; let lhs_ty = lhs.get_type(ctx); - let rhs_ty = rhs.get_type(ctx); let lhs_width = lhs_ty .deref(ctx) .downcast_ref::() @@ -495,16 +502,84 @@ where })? .width(); - let rhs_casted = if lhs_ty != rhs_ty { - let rhs_width = rhs_ty - .deref(ctx) - .downcast_ref::() - .ok_or_else(|| { - pliron::input_error!(op.deref(ctx).loc(), "Shift amount must be integer type") - })? - .width(); + let lhs = normalize_shift_value(ctx, rewriter, lhs, lhs_ty, lhs_width); + let lhs_ty = lhs.get_type(ctx); + + let rhs_casted = normalize_shift_amount(ctx, rewriter, op, rhs, lhs_ty, lhs_width)?; + let rhs_masked = mask_shift_amount(ctx, rewriter, rhs_casted, lhs_ty, lhs_width); + let llvm_op = builder(ctx, lhs, rhs_masked); + rewriter.insert_operation(ctx, llvm_op); + rewriter.replace_operation(ctx, op, llvm_op); + Ok(()) +} + +/// The shift value on the dialect's canonical integer representation. +/// +/// Every op in the chain — the count mask's `llvm.and` and the shift itself — +/// is an `IntBinArithOp`, and pliron-llvm rejects a represented (`uiN`/`siN`) +/// operand of one outright ("Integer binary arithmetic Op can only have +/// signless integer result/operand type"). Carrying a represented value into +/// the mask therefore cannot produce a legal module, whatever the count looks +/// like; the value is unified here first. +/// +/// The change is a same-width representation change of the same bits, and the +/// exporter prints every `IntegerType` as `i{width}` either way, so this is a +/// no-op in the emitted LLVM IR. +fn normalize_shift_value( + ctx: &mut Context, + rewriter: &mut DialectConversionRewriter, + lhs: Value, + lhs_ty: pliron::r#type::TypeHandle, + lhs_width: u32, +) -> Value { + let signless: pliron::r#type::TypeHandle = + IntegerType::get(ctx, lhs_width, Signedness::Signless).into(); + if lhs_ty == signless { + return lhs; + } + let cast_op = llvm::BitcastOp::new(ctx, lhs, signless).get_operation(); + rewriter.insert_operation(ctx, cast_op); + cast_op.deref(ctx).get_result(0) +} + +/// Bring a shift amount onto the shift value's integer representation. +/// +/// The width relation decides the width-changing cases: a narrower amount is +/// zero-extended, a wider one is truncated. An *equal* width with a different +/// representation is neither — it is a representation change of the same bits, +/// and it must not be expressed as a truncation: +/// +/// - LLVM `trunc` requires a strictly smaller result, so `trunc i32 -> ui32` is +/// not a legal instruction and fails module verification (#1328); +/// - handing the amount over unchanged is not enough either, because the count +/// mask and the shift itself are `SameOperandsAndResultType` ops, so the whole +/// chain has to agree on one representation. +/// +/// The same-width case therefore uses the same-width representation change the +/// cast lowering already selects for `IntToInt (same width)`. +fn normalize_shift_amount( + ctx: &mut Context, + rewriter: &mut DialectConversionRewriter, + op: Ptr, + rhs: Value, + lhs_ty: pliron::r#type::TypeHandle, + lhs_width: u32, +) -> Result { + let rhs_ty = rhs.get_type(ctx); + if rhs_ty == lhs_ty { + return Ok(rhs); + } + + let rhs_width = rhs_ty + .deref(ctx) + .downcast_ref::() + .ok_or_else(|| { + pliron::input_error!(op.deref(ctx).loc(), "Shift amount must be integer type") + })? + .width(); - let cast_op = if lhs_width > rhs_width { + let cast_op = match lhs_width.cmp(&rhs_width) { + std::cmp::Ordering::Greater => { let zext = llvm::ZExtOp::new(ctx, rhs, lhs_ty); let nneg_key: pliron::identifier::Identifier = "llvm_nneg_flag".try_into().unwrap(); zext.get_operation() @@ -512,20 +587,12 @@ where .attributes .set(nneg_key, pliron::builtin::attributes::BoolAttr::new(false)); zext.get_operation() - } else { - llvm::TruncOp::new(ctx, rhs, lhs_ty).get_operation() - }; - rewriter.insert_operation(ctx, cast_op); - cast_op.deref(ctx).get_result(0) - } else { - rhs + } + std::cmp::Ordering::Less => llvm::TruncOp::new(ctx, rhs, lhs_ty).get_operation(), + std::cmp::Ordering::Equal => llvm::BitcastOp::new(ctx, rhs, lhs_ty).get_operation(), }; - - let rhs_masked = mask_shift_amount(ctx, rewriter, rhs_casted, lhs_ty, lhs_width); - let llvm_op = builder(ctx, lhs, rhs_masked); - rewriter.insert_operation(ctx, llvm_op); - rewriter.replace_operation(ctx, op, llvm_op); - Ok(()) + rewriter.insert_operation(ctx, cast_op); + Ok(cast_op.deref(ctx).get_result(0)) } fn mask_shift_amount( @@ -538,6 +605,8 @@ fn mask_shift_amount( use pliron::utils::apint::APInt; use std::num::NonZeroUsize; + // The mask, the `and`, and the shift all live on the signless + // representation `normalize_shift_value` established. let mask_ty = IntegerType::get(ctx, lhs_width, Signedness::Signless); let mask_attr = pliron::builtin::attributes::IntegerAttr::new( mask_ty, @@ -1008,4 +1077,375 @@ mod tests { assert_eq!(fadd.fast_math_flags(&ctx).0, FastmathFlags::empty()); assert_eq!(fsub.fast_math_flags(&ctx).0, FastmathFlags::empty()); } + + // ======================================================================== + // Direct converter coverage for shift representation handling. + // + // The source pipeline normalizes every integer value through `convert_type` + // before a shift is converted, so the representation arms of the count + // normalization cannot be reached from Rust source on current main. These + // tests build the operands at that boundary by hand and call the production + // path directly, which is the only way to execute those arms. + // ======================================================================== + + use pliron::builtin::ops::ConstantOp; + use pliron::irbuild::inserter::Inserter; + use pliron::linked_list::ContainsLinkedList; + use pliron::utils::apint::APInt; + use pliron::value::Value; + use std::num::NonZeroUsize; + + /// A `builtin.constant` defining a value of the requested representation. + fn represented_operand( + ctx: &mut Context, + block: Ptr, + width: u32, + signedness: Signedness, + ) -> (TypeHandle, Value) { + let ty = IntegerType::get(ctx, width, signedness); + let attr = IntegerAttr::new( + ty, + APInt::from_u32(3, NonZeroUsize::new(width as usize).unwrap()), + ); + let op = ConstantOp::new(ctx, Box::new(attr)).get_operation(); + op.insert_at_back(block, ctx); + (ty.into(), op.deref(ctx).get_result(0)) + } + + /// Run the production shift conversion with a real insertion point and hand + /// back everything that landed in the block. + fn convert_with_real_rewriter( + ctx: &mut Context, + block: Ptr, + op: Ptr, + right: bool, + ) -> Vec> { + let mut rewriter = DialectConversionRewriter::default(); + rewriter.set_insertion_point_to_block_end(block); + // The production entry points, not a hand-built builder: `convert_shl` + // also attaches the integer-overflow flags the dialect requires. + if right { + convert_shr(ctx, &mut rewriter, op, &OperandsInfo::default()) + .expect("shift conversion failed"); + } else { + convert_shl(ctx, &mut rewriter, op, &OperandsInfo::default()) + .expect("shift conversion failed"); + } + block.deref(ctx).iter(ctx).collect() + } + + fn ops_of(ctx: &Context, ops: &[Ptr]) -> Vec> { + ops.iter() + .filter(|op| Operation::get_op::(**op, ctx).is_some()) + .copied() + .collect() + } + + fn operand_types(ctx: &Context, op: Ptr) -> Vec { + op.deref(ctx).operands().map(|v| v.get_type(ctx)).collect() + } + + fn result_types(ctx: &Context, op: Ptr) -> Vec { + (0..op.deref(ctx).get_num_results()) + .map(|i| op.deref(ctx).get_result(i).get_type(ctx)) + .collect() + } + + /// Every op the conversion produced must satisfy its own verifier. This is + /// the assertion that matters: pliron-llvm rejects a represented operand of + /// an `IntBinArithOp`, so a legal chain is signless end to end. + fn verify_all(ctx: &Context, ops: &[Ptr]) { + use pliron::common_traits::Verify; + for op in ops { + op.deref(ctx) + .verify(ctx) + .unwrap_or_else(|e| panic!("lowered op failed verification: {e}")); + } + } + + /// Shift by an operand of `count_*` representation and assert the whole + /// chain: the expected value unification, the expected count cast, a mask + /// carrying `width - 1` on the signless representation, and a shift whose + /// operands and result agree. + fn check_shift_chain( + value_width: u32, + value_signedness: Signedness, + count_width: u32, + count_signedness: Signedness, + value_unified: bool, + count_cast: &str, + ) { + let mut ctx = make_ctx(); + let (_, block) = build_kernel(&mut ctx, vec![], vec![]); + + let (value_ty, value) = represented_operand(&mut ctx, block, value_width, value_signedness); + let (_, count) = represented_operand(&mut ctx, block, count_width, count_signedness); + + let shift = Operation::new( + &mut ctx, + mir::MirShlOp::get_concrete_op_info(), + vec![value_ty], + vec![value, count], + vec![], + 0, + ); + shift.insert_at_back(block, &ctx); + + let ops = convert_with_real_rewriter(&mut ctx, block, shift, false); + verify_all(&ctx, &ops); + + let case = format!( + "value {value_width}/{value_signedness:?}, count {count_width}/{count_signedness:?}" + ); + + // A representation change on the value is one bitcast in the block. The + // count's own cast is checked separately, so the two are never confused. + let bitcasts = ops_of::(&ctx, &ops); + match count_cast { + "bitcast" => { + assert_eq!( + bitcasts.len(), + 1, + "{case}: exactly one representation change" + ); + assert_eq!( + ops_of::(&ctx, &ops).len(), + 0, + "{case}: no truncation" + ); + assert_eq!( + ops_of::(&ctx, &ops).len(), + 0, + "{case}: no widening" + ); + } + "zext" => { + assert_eq!( + ops_of::(&ctx, &ops).len(), + 1, + "{case}: one zero extension" + ); + assert_eq!(ops_of::(&ctx, &ops).len(), 0, "{case}"); + assert_eq!(bitcasts.len(), usize::from(value_unified), "{case}"); + } + "trunc" => { + assert_eq!( + ops_of::(&ctx, &ops).len(), + 1, + "{case}: one truncation" + ); + assert_eq!(ops_of::(&ctx, &ops).len(), 0, "{case}"); + assert_eq!(bitcasts.len(), usize::from(value_unified), "{case}"); + } + "none" => { + assert_eq!(ops_of::(&ctx, &ops).len(), 0, "{case}"); + assert_eq!(ops_of::(&ctx, &ops).len(), 0, "{case}"); + assert_eq!( + bitcasts.len(), + usize::from(value_unified), + "{case}: only the value is unified, if at all" + ); + } + other => panic!("unknown count cast {other}"), + } + + let expected_ty: TypeHandle = + IntegerType::get(&mut ctx, value_width, Signedness::Signless).into(); + + let and = ops_of::(&ctx, &ops); + assert_eq!(and.len(), 1, "{case}: one count mask"); + let and_op = and[0]; + let and_operands = operand_types(&ctx, and_op); + assert_eq!(and_operands.len(), 2, "{case}: mask has two operands"); + assert_eq!( + and_operands[0], and_operands[1], + "{case}: and operand types agree" + ); + assert_eq!( + and_operands[0], expected_ty, + "{case}: and works on the value width" + ); + assert_eq!( + result_types(&ctx, and_op), + vec![expected_ty], + "{case}: and result" + ); + + let mask_operand = and_op.deref(&ctx).get_operand(1); + let constant = mask_operand + .defining_op() + .and_then(|op| Operation::get_op::(op, &ctx)) + .expect("mask operand is a constant"); + let attr = constant + .get_attr_builtin_constant_value(&ctx) + .expect("mask constant carries a value"); + let integer = (&**attr as &dyn pliron::attribute::Attribute) + .downcast_ref::() + .expect("mask constant is an integer"); + assert_eq!( + integer.value(), + APInt::from_u32( + value_width - 1, + NonZeroUsize::new(value_width as usize).unwrap() + ), + "{case}: mask value is bit_width - 1" + ); + assert_eq!( + integer.get_type().deref(&ctx).signedness(), + Signedness::Signless, + "{case}: mask carries the signless representation" + ); + + let shl = ops_of::(&ctx, &ops); + assert_eq!(shl.len(), 1, "{case}: one shift"); + let shl_op = shl[0]; + let shl_operands = operand_types(&ctx, shl_op); + assert_eq!( + shl_operands[0], shl_operands[1], + "{case}: shift operand types agree" + ); + assert_eq!(shl_operands[0], expected_ty, "{case}: shift is signless"); + assert_eq!( + shl_op.deref(&ctx).get_operand(1), + and_op.deref(&ctx).get_result(0), + "{case}: the shift consumes the masked count" + ); + assert_eq!( + result_types(&ctx, shl_op), + vec![expected_ty], + "{case}: shift result" + ); + } + + /// #1328 as reported: an unsigned-32 value shifted by a signless-32 count. + /// The value is unified to the shift's canonical representation; the count, + /// already signless, is left alone. A truncation here is the reported bug. + #[test] + fn represented_value_is_unified_to_signless() { + check_shift_chain( + 32, + Signedness::Unsigned, + 32, + Signedness::Signless, + true, + "none", + ); + } + + /// The mirror case: a signless value with a represented count. The value + /// needs nothing, and the count gets the same-width representation change. + #[test] + fn represented_count_is_unified_to_signless() { + check_shift_chain( + 32, + Signedness::Signless, + 32, + Signedness::Unsigned, + false, + "bitcast", + ); + } + + /// A signed-32 value against a signless count: same-width representation + /// change on the value, no widening or narrowing anywhere. + #[test] + fn signed_value_is_unified_without_a_width_change() { + check_shift_chain( + 32, + Signedness::Signed, + 32, + Signedness::Signless, + true, + "none", + ); + } + + /// Identical representations introduce no cast at all, but the count is + /// still masked. + #[test] + fn identical_representations_introduce_no_cast() { + check_shift_chain( + 32, + Signedness::Signless, + 32, + Signedness::Signless, + false, + "none", + ); + } + + /// A narrower count is widened, and the result lands on the value's type. + #[test] + fn narrower_count_is_zero_extended() { + check_shift_chain( + 32, + Signedness::Signless, + 8, + Signedness::Signless, + false, + "zext", + ); + } + + /// A wider count is truncated, and the mask follows the value's width. + #[test] + fn wider_count_is_truncated() { + check_shift_chain( + 8, + Signedness::Signless, + 32, + Signedness::Signless, + false, + "trunc", + ); + } + + /// Right shift takes its signedness from the original MIR value, not from + /// the signless type it lowers to. + #[test] + fn right_shift_signedness_comes_from_the_original_value() { + for (signedness, expects_arithmetic) in + [(Signedness::Signed, true), (Signedness::Unsigned, false)] + { + let mut ctx = make_ctx(); + let (_, block) = build_kernel(&mut ctx, vec![], vec![]); + let (value_ty, value) = represented_operand(&mut ctx, block, 32, signedness); + let (_, count) = represented_operand(&mut ctx, block, 32, Signedness::Signless); + + let shift = Operation::new( + &mut ctx, + mir::MirShrOp::get_concrete_op_info(), + vec![value_ty], + vec![value, count], + vec![], + 0, + ); + shift.insert_at_back(block, &ctx); + + let ops = convert_with_real_rewriter(&mut ctx, block, shift, true); + verify_all(&ctx, &ops); + + let case = format!("{signedness:?} value"); + assert_eq!( + ops_of::(&ctx, &ops).len(), + usize::from(expects_arithmetic), + "{case}: arithmetic form" + ); + assert_eq!( + ops_of::(&ctx, &ops).len(), + usize::from(!expects_arithmetic), + "{case}: logical form" + ); + assert_eq!( + ops_of::(&ctx, &ops).len(), + 1, + "{case}: count masked" + ); + assert_eq!( + ops_of::(&ctx, &ops).len(), + 0, + "{case}: no truncation" + ); + } + } } diff --git a/cuda-oxide/crates/mir-lower/tests/lowering_test/main.rs b/cuda-oxide/crates/mir-lower/tests/lowering_test/main.rs index 09ab58e43a..1e93617602 100644 --- a/cuda-oxide/crates/mir-lower/tests/lowering_test/main.rs +++ b/cuda-oxide/crates/mir-lower/tests/lowering_test/main.rs @@ -17,6 +17,7 @@ mod inline_ptx; mod math_conversions; mod matrix_memory; mod mma; +mod shift_representation; mod sregs_and_warp; mod tma; mod wgmma_lowering; diff --git a/cuda-oxide/crates/mir-lower/tests/lowering_test/shift_representation.rs b/cuda-oxide/crates/mir-lower/tests/lowering_test/shift_representation.rs new file mode 100644 index 0000000000..2923a4aa99 --- /dev/null +++ b/cuda-oxide/crates/mir-lower/tests/lowering_test/shift_representation.rs @@ -0,0 +1,434 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +//! Integer shift lowering: the count-width conversion, the `bit_width - 1` +//! count mask, and the operand-type agreement `llvm.shl` / `llvm.lshr` / +//! `llvm.ashr` require. +//! +//! #1328 was an invalid same-width `llvm.trunc` on the shift count: a +//! `redux.sync` result reached the shift as `ui32` while the literal count was +//! signless, so `convert_shift` took its "types differ, so cast the count" +//! path and narrowed a 32-bit value to 32 bits. LLVM rejects that (`trunc` +//! needs a strictly smaller result) and the whole device module failed +//! verification. +//! +//! Reduction results are signless now, which is what keeps that state out of +//! the current pipeline, but the shift shape itself had no regression: the +//! redux tests cover returning a reduction directly, and no shift test runs +//! an intrinsic result into a shift. `redux_max_u32_shift_lowers` fails on the +//! tree the report was filed against and passes here. +//! +//! Every case asserts the whole chain rather than the first op: the count +//! conversion, the mask, and the shift all have to agree on one type, because +//! `llvm.and` and the shift itself are `SameOperandsAndResultType` ops. +//! Repairing only the count cast leaves the mask mismatched. + +use dialect_mir::ops as mir; +use dialect_nvvm::ops as nvvm; +use llvm_export::ops as llvm; +use pliron::basic_block::BasicBlock; +use pliron::builtin::attributes::{IntegerAttr, TypeAttr}; +use pliron::builtin::op_interfaces::SymbolOpInterface; +use pliron::builtin::ops::ModuleOp; +use pliron::builtin::types::{FunctionType, IntegerType, Signedness}; +use pliron::common_traits::Verify; +use pliron::context::{Context, Ptr}; +use pliron::linked_list::ContainsLinkedList; +use pliron::op::Op; +use pliron::operation::Operation; +use pliron::printable::Printable; +use pliron::r#type::{TypeHandle, Typed}; +use pliron::utils::apint::APInt; +use pliron::value::Value; +use std::num::NonZeroUsize; + +use crate::common::{lowered_kernel_body, make_test_ctx}; + +const KERNEL: &str = "kernel_func"; + +fn integer(ctx: &mut Context, width: u32, signedness: Signedness) -> TypeHandle { + IntegerType::get(ctx, width, signedness).into() +} + +/// A MIR function `kernel_func(args) -> (ret)` with one entry block. +fn build_returning_kernel( + ctx: &mut Context, + arg_tys: Vec, + ret_tys: Vec, +) -> (Ptr, Ptr) { + let module = ModuleOp::new(ctx, "shift_representation".try_into().unwrap()); + let module_ptr = module.get_operation(); + + let func_ty = FunctionType::get(ctx, arg_tys.clone(), ret_tys); + let func_ptr = Operation::new( + ctx, + mir::MirFuncOp::get_concrete_op_info(), + vec![], + vec![], + vec![], + 1, + ); + let func = mir::MirFuncOp::new(ctx, func_ptr, TypeAttr::new(func_ty.into())); + func.set_symbol_name(ctx, KERNEL.try_into().unwrap()); + + let region = func.get_operation().deref(ctx).get_region(0); + let entry = BasicBlock::new(ctx, None, arg_tys); + entry.insert_at_back(region, ctx); + + let module_block = module_ptr + .deref(ctx) + .get_region(0) + .deref(ctx) + .iter(ctx) + .next() + .unwrap(); + func.get_operation().insert_at_back(module_block, ctx); + + (module_ptr, entry) +} + +fn mir_integer_constant( + ctx: &mut Context, + block: Ptr, + width: u32, + signedness: Signedness, + value: u32, +) -> Value { + let ty = IntegerType::get(ctx, width, signedness); + let constant = Operation::new( + ctx, + mir::MirConstantOp::get_concrete_op_info(), + vec![ty.into()], + vec![], + vec![], + 0, + ); + mir::MirConstantOp::new(constant).set_attr_value( + ctx, + IntegerAttr::new( + ty, + APInt::from_u32(value, NonZeroUsize::new(width as usize).unwrap()), + ), + ); + constant.insert_at_back(block, ctx); + constant.deref(ctx).get_result(0) +} + +fn insert_shift(ctx: &mut Context, block: Ptr, op: Ptr) -> Value { + op.insert_at_back(block, ctx); + op.deref(ctx).get_result(0) +} + +/// Build a `mir.shl` whose result carries `value_ty`, as Rust's `Shl` impls do. +fn mir_shl( + ctx: &mut Context, + block: Ptr, + value_ty: TypeHandle, + lhs: Value, + rhs: Value, +) -> Value { + let op = Operation::new( + ctx, + mir::MirShlOp::get_concrete_op_info(), + vec![value_ty], + vec![lhs, rhs], + vec![], + 0, + ); + insert_shift(ctx, block, op) +} + +fn mir_shr( + ctx: &mut Context, + block: Ptr, + value_ty: TypeHandle, + lhs: Value, + rhs: Value, +) -> Value { + let op = Operation::new( + ctx, + mir::MirShrOp::get_concrete_op_info(), + vec![value_ty], + vec![lhs, rhs], + vec![], + 0, + ); + insert_shift(ctx, block, op) +} + +fn mir_return(ctx: &mut Context, block: Ptr, values: Vec) { + let op = Operation::new( + ctx, + mir::MirReturnOp::get_concrete_op_info(), + vec![], + values, + vec![], + 0, + ); + op.insert_at_back(block, ctx); +} + +/// Lower, verify, and hand back the kernel body for inspection. +fn lower_and_verify(ctx: &mut Context, module_ptr: Ptr) -> Vec> { + mir_lower::lower_mir_to_llvm(ctx, module_ptr).expect("shift lowering failed"); + module_ptr + .deref(ctx) + .verify(ctx) + .expect("lowered module must verify"); + lowered_kernel_body(ctx, module_ptr) +} + +fn count_of(ctx: &Context, body: &[Ptr]) -> usize { + body.iter() + .filter(|op| Operation::get_op::(**op, ctx).is_some()) + .count() +} + +fn find(ctx: &Context, body: &[Ptr]) -> T { + body.iter() + .find_map(|op| Operation::get_op::(*op, ctx)) + .expect("expected op in lowered kernel body") +} + +fn operand_types(ctx: &Context, op: Ptr) -> Vec { + op.deref(ctx).operands().map(|v| v.get_type(ctx)).collect() +} + +fn result_types(ctx: &Context, op: Ptr) -> Vec { + (0..op.deref(ctx).get_num_results()) + .map(|i| op.deref(ctx).get_result(i).get_type(ctx)) + .collect() +} + +/// The reported #1328 shape: an unsigned 32-bit warp reduction shifted left by +/// a literal. The reduction result and the count share a width, so the count +/// must reach `llvm.shl` unconverted; a same-width `trunc` here is the bug. +#[test] +fn redux_max_u32_shift_lowers() { + let mut ctx = make_test_ctx(); + let u32_ty = integer(&mut ctx, 32, Signedness::Unsigned); + + let (module_ptr, block) = build_returning_kernel(&mut ctx, vec![u32_ty, u32_ty], vec![u32_ty]); + let mask = block.deref(&ctx).get_argument(0); + let value = block.deref(&ctx).get_argument(1); + + let redux = nvvm::ReduxSyncUmaxOp::build(&mut ctx, mask, value); + redux.insert_at_back(block, &ctx); + let maximum = redux.deref(&ctx).get_result(0); + + let count = mir_integer_constant(&mut ctx, block, 32, Signedness::Unsigned, 8); + let shifted = mir_shl(&mut ctx, block, u32_ty, maximum, count); + mir_return(&mut ctx, block, vec![shifted]); + + let body = lower_and_verify(&mut ctx, module_ptr); + + assert_eq!( + count_of::(&ctx, &body), + 0, + "no count narrowing" + ); + assert_eq!( + count_of::(&ctx, &body), + 0, + "no count widening" + ); + + let shl = find::(&ctx, &body); + let shl_op = shl.get_operation(); + let operands = operand_types(&ctx, shl_op); + assert_eq!(operands.len(), 2); + assert_eq!( + operands[0], + operands[1], + "llvm.shl requires one operand type, got {} and {}", + operands[0].disp(&ctx), + operands[1].disp(&ctx) + ); + assert_eq!( + result_types(&ctx, shl_op), + vec![operands[0]], + "llvm.shl result must match its operands" + ); +} + +/// The count mask is part of the same chain: `llvm.and` is a +/// `SameOperandsAndResultType` op, so a mask built in a different +/// representation than the count fails verification just as the shift would. +#[test] +fn shift_count_mask_shares_the_shift_operand_type() { + let mut ctx = make_test_ctx(); + let u32_ty = integer(&mut ctx, 32, Signedness::Unsigned); + + let (module_ptr, block) = build_returning_kernel(&mut ctx, vec![u32_ty, u32_ty], vec![u32_ty]); + let mask = block.deref(&ctx).get_argument(0); + let value = block.deref(&ctx).get_argument(1); + + let redux = nvvm::ReduxSyncUmaxOp::build(&mut ctx, mask, value); + redux.insert_at_back(block, &ctx); + let maximum = redux.deref(&ctx).get_result(0); + + let count = mir_integer_constant(&mut ctx, block, 32, Signedness::Unsigned, 8); + let shifted = mir_shl(&mut ctx, block, u32_ty, maximum, count); + mir_return(&mut ctx, block, vec![shifted]); + + let body = lower_and_verify(&mut ctx, module_ptr); + + let and = find::(&ctx, &body); + let and_op = and.get_operation(); + let and_operands = operand_types(&ctx, and_op); + assert_eq!(and_operands.len(), 2, "one mask and per shift"); + assert_eq!( + and_operands[0], and_operands[1], + "llvm.and requires one operand type" + ); + assert_eq!(result_types(&ctx, and_op), vec![and_operands[0]]); + + // The masked count is what the shift consumes, so both chains agree on one + // type end to end. + let shl = find::(&ctx, &body); + assert_eq!( + shl.get_operation().deref(&ctx).get_operand(1), + and_op.deref(&ctx).get_result(0), + "the shift consumes the masked count" + ); + assert_eq!(operand_types(&ctx, shl.get_operation())[1], and_operands[0]); +} + +/// A count narrower than the value widens; a count wider than the value +/// narrows. Both land the count on the value's type without a representation +/// change, and the mask follows. +#[test] +fn shift_count_widths_convert_in_the_documented_direction() { + for (value_width, count_width, expects_zext) in [(32u32, 8u32, true), (8, 32, false)] { + let mut ctx = make_test_ctx(); + let value_ty = integer(&mut ctx, value_width, Signedness::Unsigned); + let count_ty = integer(&mut ctx, count_width, Signedness::Unsigned); + + let (module_ptr, block) = + build_returning_kernel(&mut ctx, vec![value_ty, count_ty], vec![value_ty]); + let value = block.deref(&ctx).get_argument(0); + let count = block.deref(&ctx).get_argument(1); + + let shifted = mir_shl(&mut ctx, block, value_ty, value, count); + mir_return(&mut ctx, block, vec![shifted]); + + let body = lower_and_verify(&mut ctx, module_ptr); + + let case = format!("value {value_width}, count {count_width}"); + assert_eq!( + count_of::(&ctx, &body), + usize::from(expects_zext), + "{case}: zero-extension count" + ); + assert_eq!( + count_of::(&ctx, &body), + usize::from(!expects_zext), + "{case}: truncating count" + ); + assert_eq!( + count_of::(&ctx, &body), + 0, + "{case}: a width change is never a representation change" + ); + + // The count lands on the value's width in the lowered signless + // representation, which is the type the whole chain has to share. + let shl = find::(&ctx, &body); + let operands = operand_types(&ctx, shl.get_operation()); + assert_eq!(operands[0], operands[1], "{case}: one shift operand type"); + let lowered_width = operands[0] + .deref(&ctx) + .downcast_ref::() + .expect("shift operands are integers") + .width(); + assert_eq!( + lowered_width, value_width, + "{case}: count lands on value width" + ); + } +} + +/// Signedness of a right shift is a property of the original MIR operation, +/// not of the lowered signless type: a signed value takes `ashr`, an unsigned +/// one takes `lshr`, and neither skips the count mask. +#[test] +fn right_shift_selects_arithmetic_or_logical_from_mir_signedness() { + for (signedness, expects_arithmetic) in + [(Signedness::Signed, true), (Signedness::Unsigned, false)] + { + let mut ctx = make_test_ctx(); + let value_ty = integer(&mut ctx, 32, signedness); + + let (module_ptr, block) = + build_returning_kernel(&mut ctx, vec![value_ty, value_ty], vec![value_ty]); + let value = block.deref(&ctx).get_argument(0); + let count = block.deref(&ctx).get_argument(1); + + let shifted = mir_shr(&mut ctx, block, value_ty, value, count); + mir_return(&mut ctx, block, vec![shifted]); + + let body = lower_and_verify(&mut ctx, module_ptr); + + let case = format!("{signedness:?} value"); + assert_eq!( + count_of::(&ctx, &body), + usize::from(expects_arithmetic), + "{case}: arithmetic form" + ); + assert_eq!( + count_of::(&ctx, &body), + usize::from(!expects_arithmetic), + "{case}: logical form" + ); + assert_eq!(count_of::(&ctx, &body), 0, "{case}"); + assert_eq!(count_of::(&ctx, &body), 0, "{case}"); + } +} + +/// An equal-width count is left alone (no cast at all) but is still masked with +/// `bit_width - 1`, and the shift consumes the mask rather than the raw count. +#[test] +fn matching_count_width_is_masked_without_a_cast() { + let mut ctx = make_test_ctx(); + let u8_ty = integer(&mut ctx, 8, Signedness::Unsigned); + + let (module_ptr, block) = build_returning_kernel(&mut ctx, vec![u8_ty, u8_ty], vec![u8_ty]); + let value = block.deref(&ctx).get_argument(0); + let count = block.deref(&ctx).get_argument(1); + + let shifted = mir_shl(&mut ctx, block, u8_ty, value, count); + mir_return(&mut ctx, block, vec![shifted]); + + let body = lower_and_verify(&mut ctx, module_ptr); + + assert_eq!(count_of::(&ctx, &body), 0); + assert_eq!(count_of::(&ctx, &body), 0); + assert_eq!(count_of::(&ctx, &body), 0); + + let and = find::(&ctx, &body); + let and_op = and.get_operation(); + let mask_operand = and_op.deref(&ctx).get_operand(1); + let constant = mask_operand + .defining_op() + .and_then(|op| Operation::get_op::(op, &ctx)) + .expect("the mask operand is a constant"); + let attr = constant + .get_attr_builtin_constant_value(&ctx) + .expect("the mask constant carries a value"); + let integer_attr = (&**attr as &dyn pliron::attribute::Attribute) + .downcast_ref::() + .expect("the mask constant is an integer"); + assert_eq!( + integer_attr.value(), + APInt::from_u32(7, NonZeroUsize::new(8).unwrap()), + "the mask is bit_width - 1" + ); + + let shl = find::(&ctx, &body); + assert_eq!( + shl.get_operation().deref(&ctx).get_operand(1), + and_op.deref(&ctx).get_result(0) + ); +} diff --git a/cuda-oxide/crates/rustc-codegen-cuda/examples/redux_shift_regression/Cargo.lock b/cuda-oxide/crates/rustc-codegen-cuda/examples/redux_shift_regression/Cargo.lock new file mode 100644 index 0000000000..c23ab7db71 --- /dev/null +++ b/cuda-oxide/crates/rustc-codegen-cuda/examples/redux_shift_regression/Cargo.lock @@ -0,0 +1,724 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 4 + +[[package]] +name = "aho-corasick" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ddd31a130427c27518df266943a5308ed92d4b226cc639f5a8f1002816174301" +dependencies = [ + "memchr", +] + +[[package]] +name = "anyhow" +version = "1.0.102" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7f202df86484c868dbad7eaa557ef785d5c66295e41b460ef922eca0723b842c" + +[[package]] +name = "bindgen" +version = "0.69.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "271383c67ccabffb7381723dea0672a673f292304fcb45c01cc648c7a8d58088" +dependencies = [ + "bitflags", + "cexpr", + "clang-sys", + "itertools", + "lazy_static", + "lazycell", + "log", + "prettyplease", + "proc-macro2", + "quote", + "regex", + "rustc-hash", + "shlex", + "syn", + "which", +] + +[[package]] +name = "bitflags" +version = "2.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b4388bee8683e3d04af747c73422af53102d2bd24d9eadb6cbc100baef4b43f8" + +[[package]] +name = "block-buffer" +version = "0.10.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71" +dependencies = [ + "generic-array", +] + +[[package]] +name = "cexpr" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6fac387a98bb7c37292057cffc56d62ecb629900026402633ae9160df93a8766" +dependencies = [ + "nom", +] + +[[package]] +name = "cfg-if" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" + +[[package]] +name = "clang-sys" +version = "1.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b023947811758c97c59bf9d1c188fd619ad4718dcaa767947df1cadb14f39f4" +dependencies = [ + "glob", + "libc", + "libloading", +] + +[[package]] +name = "cpufeatures" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280" +dependencies = [ + "libc", +] + +[[package]] +name = "crc32fast" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9481c1c90cbf2ac953f07c8d4a58aa3945c425b7185c9154d67a65e4230da511" +dependencies = [ + "cfg-if", +] + +[[package]] +name = "crunchy" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5" + +[[package]] +name = "crypto-common" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" +dependencies = [ + "generic-array", + "typenum", +] + +[[package]] +name = "cuda-artifact-finalizer" +version = "0.2.1" +dependencies = [ + "libnvvm-sys", + "nvjitlink-sys", + "serde", + "sha2", + "thiserror", +] + +[[package]] +name = "cuda-bindings" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0f42e2c99c56bf2eef9d569435bfd72a478dfaea3f178b8e5bae891875bbbe2e" +dependencies = [ + "bindgen", + "libloading", + "prettyplease", + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "cuda-core" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "84bdb592a2844148fb13bdfbbab158879f56916bde1d1a1b6580ca27f23850cd" +dependencies = [ + "anyhow", + "cuda-bindings", + "cuda-core-derive", + "half", + "oxide-artifacts", +] + +[[package]] +name = "cuda-core-derive" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8a48da7d10c1078f7fbb87d7176e9476affadea148babf31946e348073872d44" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "cuda-device" +version = "0.2.1" +dependencies = [ + "cuda-macros", +] + +[[package]] +name = "cuda-host" +version = "0.2.1" +dependencies = [ + "cuda-artifact-finalizer", + "cuda-core", + "cuda-macros", + "half", + "libc", + "oxide-artifacts", + "ptx-parse", + "sha2", + "thiserror", +] + +[[package]] +name = "cuda-macros" +version = "0.2.1" +dependencies = [ + "proc-macro2", + "quote", + "reserved-oxide-symbols", + "syn", +] + +[[package]] +name = "cuda-target-spec" +version = "0.1.0" + +[[package]] +name = "digest" +version = "0.10.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" +dependencies = [ + "block-buffer", + "crypto-common", +] + +[[package]] +name = "either" +version = "1.16.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91622ff5e7162018101f2fea40d6ebf4a78bbe5a49736a2020649edf9693679e" + +[[package]] +name = "equivalent" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" + +[[package]] +name = "errno" +version = "0.3.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" +dependencies = [ + "libc", + "windows-sys 0.61.2", +] + +[[package]] +name = "foldhash" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d9c4f5dac5e15c24eb999c26181a6ca40b39fe946cbe4c263c7209467bc83af2" + +[[package]] +name = "generic-array" +version = "0.14.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" +dependencies = [ + "typenum", + "version_check", +] + +[[package]] +name = "glob" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0cc23270f6e1808e30a928bdc84dea0b9b4136a8bc82338574f23baf47bbd280" + +[[package]] +name = "half" +version = "2.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ea2d84b969582b4b1864a92dc5d27cd2b77b622a8d79306834f1be5ba20d84b" +dependencies = [ + "cfg-if", + "crunchy", + "zerocopy", +] + +[[package]] +name = "hashbrown" +version = "0.15.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9229cfe53dfd69f0609a49f65461bd93001ea1ef889cd5529dd176593f5338a1" +dependencies = [ + "foldhash", +] + +[[package]] +name = "hashbrown" +version = "0.17.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" + +[[package]] +name = "home" +version = "0.5.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cc627f471c528ff0c4a49e1d5e60450c8f6461dd6d10ba9dcd3a61d3dff7728d" +dependencies = [ + "windows-sys 0.61.2", +] + +[[package]] +name = "indexmap" +version = "2.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" +dependencies = [ + "equivalent", + "hashbrown 0.17.1", +] + +[[package]] +name = "itertools" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba291022dbbd398a455acf126c1e341954079855bc60dfdda641363bd6922569" +dependencies = [ + "either", +] + +[[package]] +name = "lazy_static" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" + +[[package]] +name = "lazycell" +version = "1.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "830d08ce1d1d941e6b30645f1a0eb5643013d835ce3779a5fc208261dbe10f55" + +[[package]] +name = "libc" +version = "0.2.186" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" + +[[package]] +name = "libloading" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d7c4b02199fee7c5d21a5ae7d8cfa79a6ef5bb2fc834d6e9058e89c825efdc55" +dependencies = [ + "cfg-if", + "windows-link", +] + +[[package]] +name = "libnvvm-sys" +version = "0.2.1" +dependencies = [ + "cuda-target-spec", + "libloading", + "thiserror", +] + +[[package]] +name = "linux-raw-sys" +version = "0.4.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d26c52dbd32dccf2d10cac7725f8eae5296885fb5703b261f7d0a0739ec807ab" + +[[package]] +name = "log" +version = "0.4.29" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5e5032e24019045c762d3c0f28f5b6b8bbf38563a65908389bf7978758920897" + +[[package]] +name = "memchr" +version = "2.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98" + +[[package]] +name = "minimal-lexical" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" + +[[package]] +name = "nom" +version = "7.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d273983c5a657a70a3e8f2a01329822f3b8c8172b73826411a55751e404a0a4a" +dependencies = [ + "memchr", + "minimal-lexical", +] + +[[package]] +name = "nvjitlink-sys" +version = "0.2.1" +dependencies = [ + "libloading", + "thiserror", +] + +[[package]] +name = "object" +version = "0.36.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "62948e14d923ea95ea2c7c86c71013138b66525b86bdc08d2dcc262bdb497b87" +dependencies = [ + "crc32fast", + "hashbrown 0.15.5", + "indexmap", + "memchr", +] + +[[package]] +name = "once_cell" +version = "1.21.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" + +[[package]] +name = "oxide-artifacts" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "74d9a941c0c43d2dceaf25f5c9ca4c9f9cceb28dfa5e4e3cfd14e476f003e7c9" +dependencies = [ + "object", +] + +[[package]] +name = "prettyplease" +version = "0.2.37" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "479ca8adacdd7ce8f1fb39ce9ecccbfe93a3f1344b3d0d97f20bc0196208f62b" +dependencies = [ + "proc-macro2", + "syn", +] + +[[package]] +name = "proc-macro2" +version = "1.0.106" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "ptx-parse" +version = "0.1.0" + +[[package]] +name = "quote" +version = "1.0.45" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41f2619966050689382d2b44f664f4bc593e129785a36d6ee376ddf37259b924" +dependencies = [ + "proc-macro2", +] + +[[package]] +name = "redux_shift_regression" +version = "0.1.0" +dependencies = [ + "cuda-core", + "cuda-device", + "cuda-host", +] + +[[package]] +name = "regex" +version = "1.12.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e10754a14b9137dd7b1e3e5b0493cc9171fdd105e0ab477f51b72e7f3ac0e276" +dependencies = [ + "aho-corasick", + "memchr", + "regex-automata", + "regex-syntax", +] + +[[package]] +name = "regex-automata" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6e1dd4122fc1595e8162618945476892eefca7b88c52820e74af6262213cae8f" +dependencies = [ + "aho-corasick", + "memchr", + "regex-syntax", +] + +[[package]] +name = "regex-syntax" +version = "0.8.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" + +[[package]] +name = "reserved-oxide-symbols" +version = "0.2.1" + +[[package]] +name = "rustc-hash" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08d43f7aa6b08d49f382cde6a7982047c3426db949b1424bc4b7ec9ae12c6ce2" + +[[package]] +name = "rustix" +version = "0.38.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fdb5bc1ae2baa591800df16c9ca78619bf65c0488b41b96ccec5d11220d8c154" +dependencies = [ + "bitflags", + "errno", + "libc", + "linux-raw-sys", + "windows-sys 0.59.0", +] + +[[package]] +name = "serde" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" +dependencies = [ + "serde_core", + "serde_derive", +] + +[[package]] +name = "serde_core" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad" +dependencies = [ + "serde_derive", +] + +[[package]] +name = "serde_derive" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "sha2" +version = "0.10.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" +dependencies = [ + "cfg-if", + "cpufeatures", + "digest", +] + +[[package]] +name = "shlex" +version = "1.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0fda2ff0d084019ba4d7c6f371c95d8fd75ce3524c3cb8fb653a3023f6323e64" + +[[package]] +name = "syn" +version = "2.0.117" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e665b8803e7b1d2a727f4023456bbbbe74da67099c585258af0ad9c5013b9b99" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "thiserror" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4288b5bcbc7920c07a1149a35cf9590a2aa808e0bc1eafaade0b80947865fbc4" +dependencies = [ + "thiserror-impl", +] + +[[package]] +name = "thiserror-impl" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "typenum" +version = "1.20.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6f5e870be6c3b371b77fe0ee0bafb859fa4964b4404c27de1d380043c4dda20" + +[[package]] +name = "unicode-ident" +version = "1.0.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" + +[[package]] +name = "version_check" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" + +[[package]] +name = "which" +version = "4.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "87ba24419a2078cd2b0f2ede2691b6c66d8e47836da3b6db8265ebad47afbfc7" +dependencies = [ + "either", + "home", + "once_cell", + "rustix", +] + +[[package]] +name = "windows-link" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" + +[[package]] +name = "windows-sys" +version = "0.59.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e38bc4d79ed67fd075bcc251a1c39b32a1776bbe92e5bef1f0bf1f8c531853b" +dependencies = [ + "windows-targets", +] + +[[package]] +name = "windows-sys" +version = "0.61.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc" +dependencies = [ + "windows-link", +] + +[[package]] +name = "windows-targets" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973" +dependencies = [ + "windows_aarch64_gnullvm", + "windows_aarch64_msvc", + "windows_i686_gnu", + "windows_i686_gnullvm", + "windows_i686_msvc", + "windows_x86_64_gnu", + "windows_x86_64_gnullvm", + "windows_x86_64_msvc", +] + +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" + +[[package]] +name = "windows_aarch64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" + +[[package]] +name = "windows_i686_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b" + +[[package]] +name = "windows_i686_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" + +[[package]] +name = "windows_i686_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" + +[[package]] +name = "windows_x86_64_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" + +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" + +[[package]] +name = "windows_x86_64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" + +[[package]] +name = "zerocopy" +version = "0.8.52" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ce1022995ff5ff5d841ad7d994facc23098cd40152f2c1d11cd607c6f530653f" +dependencies = [ + "zerocopy-derive", +] + +[[package]] +name = "zerocopy-derive" +version = "0.8.52" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ae7f38b72ec2a254e2b87ef277cf2cd4fb97cbebf944faa6f33354da0867930" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] diff --git a/cuda-oxide/crates/rustc-codegen-cuda/examples/redux_shift_regression/Cargo.toml b/cuda-oxide/crates/rustc-codegen-cuda/examples/redux_shift_regression/Cargo.toml new file mode 100644 index 0000000000..992b15da65 --- /dev/null +++ b/cuda-oxide/crates/rustc-codegen-cuda/examples/redux_shift_regression/Cargo.toml @@ -0,0 +1,24 @@ +[package] +name = "redux_shift_regression" +version = "0.1.0" +edition = "2024" +license = "Apache-2.0" + +# Mark as standalone crate (not part of parent workspace) +[workspace] + +# Regression for #1328: shifting a `redux.sync` result used to lower through an +# illegal same-width `llvm.trunc`, because the reduction result carried the +# catalog's `ui32` representation while the shift count was signless. Build and +# run with: +# cargo oxide run redux_shift_regression --arch sm_80 + +[dependencies] +# Device-side intrinsics (warp::redux_sync_max_u32, _min_i32) +cuda-device = { path = "../../../cuda-device" } + +# Host-side utilities (CudaKernel trait, etc.) +cuda-host = { path = "../../../cuda-host" } + +# Host-side CUDA runtime +cuda-core = "0.3.1" diff --git a/cuda-oxide/crates/rustc-codegen-cuda/examples/redux_shift_regression/src/main.rs b/cuda-oxide/crates/rustc-codegen-cuda/examples/redux_shift_regression/src/main.rs new file mode 100644 index 0000000000..ade625f907 --- /dev/null +++ b/cuda-oxide/crates/rustc-codegen-cuda/examples/redux_shift_regression/src/main.rs @@ -0,0 +1,298 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +//! Shifting the result of a warp reduction (#1328). +//! +//! A `redux.sync` result used to reach the shift lowering carrying the +//! intrinsic catalog's `ui32` representation while the shift count was +//! signless. `convert_shift` saw "operands differ", compared widths, found +//! them equal and emitted `llvm.trunc i32 -> ui32` — which LLVM rejects, +//! because `trunc` needs a strictly smaller result. The whole device module +//! failed verification, so a plain `maximum << 8` did not compile. +//! +//! The reduction results are translated to signless integers now, but nothing +//! covered the shape: `redux_minmax` writes a reduction straight to memory and +//! `redux_sum` consumes it without an operator in between. This example puts +//! the operators back on and checks the values on the device: +//! +//! * `<<` / `>>` after `redux.sync.max.u32` — unsigned, so the right shift +//! must stay logical even when the top bit is set; +//! * `>>` after `redux.sync.min.s32` — signed, so it must become arithmetic +//! on negative lanes; +//! * `wrapping_shl` with counts at and past the bit width, which is the +//! mask-the-count behaviour `convert_shift` implements; +//! * a multi-warp block whose warps hold different values, so a warp's +//! reduction cannot leak into its neighbour's answer unnoticed. +//! +//! Every lane writes its own result: `redux.sync` broadcasts to the lanes named +//! by the member mask, so the shifted value has to agree across all 32. +//! +//! Build and run with: +//! cargo oxide run redux_shift_regression --arch sm_80 + +use cuda_device::{DisjointSlice, kernel, thread, warp}; +use cuda_host::cuda_module; + +const FULL_MASK: u32 = 0xffff_ffff; + +// ============================================================================= +// KERNELS +// ============================================================================= +#[cuda_module] +mod kernels { + use super::*; + + /// Lane `i` shifts the unsigned maximum of its own warp left. + #[kernel] + pub fn shl_u32(input: &[u32], count: &[u32], mut out: DisjointSlice) { + let i = thread::index_1d().get() as usize; + let maximum = warp::redux_sync_max_u32(FULL_MASK, input[i]); + let shifted = maximum << (count[i] & 31); + unsafe { + *out.get_unchecked_mut(i) = shifted; + } + } + + /// The same reduction shifted right: unsigned, so the top bit shifts in as + /// zero. A signless lowering that picked `ashr` would answer differently on + /// the patterns that set bit 31. + #[kernel] + pub fn shr_u32(input: &[u32], count: &[u32], mut out: DisjointSlice) { + let i = thread::index_1d().get() as usize; + let maximum = warp::redux_sync_max_u32(FULL_MASK, input[i]); + let shifted = maximum >> (count[i] & 31); + unsafe { + *out.get_unchecked_mut(i) = shifted; + } + } + + /// A signed reduction shifted right: negative inputs make `ashr` and `lshr` + /// disagree, which is the control for the case above. + #[kernel] + pub fn shr_i32(input: &[i32], count: &[u32], mut out: DisjointSlice) { + let i = thread::index_1d().get() as usize; + let minimum = warp::redux_sync_min_i32(FULL_MASK, input[i]); + let shifted = minimum >> (count[i] & 31); + unsafe { + *out.get_unchecked_mut(i) = shifted; + } + } + + /// Counts at and past the bit width: `wrapping_shl` and the lowering both + /// mask with `bit_width - 1`, so 32 shifts by 0 and 33 by 1. + #[kernel] + pub fn wrapping_shl_u32(input: &[u32], count: &[u32], mut out: DisjointSlice) { + let i = thread::index_1d().get() as usize; + let maximum = warp::redux_sync_max_u32(FULL_MASK, input[i]); + let shifted = maximum.wrapping_shl(count[i]); + unsafe { + *out.get_unchecked_mut(i) = shifted; + } + } +} + +// ============================================================================= +// HOST +// ============================================================================= + +/// Threads per block: four warps, so the multi-warp case is the default one. +const BLOCK: u32 = 128; + +/// Fixed-seed LCG, so the "random" pattern is reproducible without a crate. +fn lcg(state: &mut u32) -> u32 { + *state = state.wrapping_mul(1_664_525).wrapping_add(1_013_904_223); + *state +} + +/// `ramp`: lane 0..31 counted up, offset by the warp, so each warp has a +/// different unsigned max. Warps that read each other's answer cannot produce +/// the expected values. +fn ramp(len: usize) -> Vec { + (0..len) + .map(|i| (i as u32 / 32) * 64 + (i % 32) as u32) + .collect() +} + +/// `extremes`: each warp carries 0, 1, `0x8000_0000` and `0xffff_ffff`, so the +/// unsigned max sets every bit and the signed min is `i32::MIN` on those lanes. +fn extremes(len: usize) -> Vec { + (0..len) + .map(|i| match i % 32 { + 0 => 0, + 1 => 1, + 2 => 0x8000_0000, + 3 => 0xffff_ffff, + lane => (lane as u32) << 8, + }) + .collect() +} + +/// `random`: fixed-seed values with no structure a constant folder could use. +fn random(len: usize) -> Vec { + let mut state = 0x1234_5678; + (0..len).map(|_| lcg(&mut state)).collect() +} + +/// Bitwise inverse of a slice, used as an output poison. +fn invert(values: &[u32]) -> Vec { + values.iter().map(|value| !value).collect() +} + +fn invert_i32(values: &[i32]) -> Vec { + values.iter().map(|value| !value).collect() +} + +fn expected_u32(input: &[u32], count: &[u32], shift: impl Fn(u32, u32) -> u32) -> Vec { + (0..input.len()) + .map(|i| { + let base = i & !31; + let maximum = input[base..base + 32].iter().copied().max().unwrap(); + shift(maximum, count[i] & 31) + }) + .collect() +} + +fn expected_i32(input: &[i32], count: &[u32], shift: impl Fn(i32, u32) -> i32) -> Vec { + (0..input.len()) + .map(|i| { + let base = i & !31; + let minimum = input[base..base + 32].iter().copied().min().unwrap(); + shift(minimum, count[i] & 31) + }) + .collect() +} + +fn counts(len: usize, values: &[u32]) -> Vec { + (0..len).map(|i| values[i % values.len()]).collect() +} + +fn main() { + use cuda_core::simt::LaunchConfig; + use cuda_core::{CudaContext, DeviceBuffer}; + + let ctx = CudaContext::new(0).expect("Failed to create CUDA context"); + let stream = ctx.default_stream(); + let module = kernels::load(&ctx).expect("Failed to load module"); + + let cfg = LaunchConfig { + grid_dim: (1, 1, 1), + block_dim: (BLOCK, 1, 1), + shared_mem_bytes: 0, + }; + + let len = BLOCK as usize; + let mut failed = false; + + let patterns: [(&str, Vec); 3] = [ + ("ramp", ramp(len)), + ("extremes", extremes(len)), + ("random", random(len)), + ]; + let in_range = counts(len, &[0, 1, 8, 31]); + let past_width = counts(len, &[32, 33, 63]); + + // ===== Test 1: left shift, and the logical right shift control ===== + println!("--- Test 1: redux.sync.max.u32 << and >> ---"); + for (name, input) in &patterns { + let input_dev = DeviceBuffer::from_host(&stream, input).unwrap(); + let count_dev = DeviceBuffer::from_host(&stream, &in_range).unwrap(); + let want_shl = expected_u32(input, &in_range, |v, c| v << c); + let want_shr = expected_u32(input, &in_range, |v, c| v >> c); + // Initialised to the bitwise inverse of the oracle: a lane the kernel + // never writes stays at a value no correct answer can equal. + let mut shl_dev = DeviceBuffer::from_host(&stream, &invert(&want_shl)).unwrap(); + let mut shr_dev = DeviceBuffer::from_host(&stream, &invert(&want_shr)).unwrap(); + + // SAFETY: launch shape matches the buffers; every lane stays in range. + unsafe { module.shl_u32(stream.as_ref(), cfg, &input_dev, &count_dev, &mut shl_dev) } + .expect("shl_u32 launch failed"); + // SAFETY: as above. + unsafe { module.shr_u32(stream.as_ref(), cfg, &input_dev, &count_dev, &mut shr_dev) } + .expect("shr_u32 launch failed"); + + let shl = shl_dev.to_host_vec(&stream).unwrap(); + let shr = shr_dev.to_host_vec(&stream).unwrap(); + + if shl == want_shl { + println!("✓ {name}: << matches host for all {len} lanes"); + } else { + println!("✗ {name}: << mismatch"); + println!(" got {:?}", &shl[..8]); + println!(" want {:?}", &want_shl[..8]); + failed = true; + } + if shr == want_shr { + println!("✓ {name}: >> is logical for all {len} lanes"); + } else { + println!("✗ {name}: >> mismatch"); + println!(" got {:?}", &shr[..8]); + println!(" want {:?}", &want_shr[..8]); + failed = true; + } + + // The logical claim is only tested if some answer has bit 31 set. + if name == &"extremes" && want_shr.iter().all(|v| v >> 31 == 0) { + panic!("the logical-shift control never set bit 31"); + } + } + + // ===== Test 2: signed right shift stays arithmetic ===== + println!("\n--- Test 2: redux.sync.min.s32 >> ---"); + { + let signed: Vec = extremes(len).iter().map(|v| *v as i32).collect(); + let want = expected_i32(&signed, &in_range, |v, c| v >> c); + let input_dev = DeviceBuffer::from_host(&stream, &signed).unwrap(); + let count_dev = DeviceBuffer::from_host(&stream, &in_range).unwrap(); + let mut out_dev = DeviceBuffer::from_host(&stream, &invert_i32(&want)).unwrap(); + + // SAFETY: launch shape matches the buffers; every lane stays in range. + unsafe { module.shr_i32(stream.as_ref(), cfg, &input_dev, &count_dev, &mut out_dev) } + .expect("shr_i32 launch failed"); + + let got = out_dev.to_host_vec(&stream).unwrap(); + if got == want { + println!("✓ signed >> is arithmetic for all {len} lanes"); + } else { + println!("✗ signed >> mismatch"); + println!(" got {:?}", &got[..8]); + println!(" want {:?}", &want[..8]); + failed = true; + } + if want.iter().all(|v| *v >= 0) { + panic!("the arithmetic-shift control never produced a negative answer"); + } + } + + // ===== Test 3: counts at and past the bit width ===== + println!("\n--- Test 3: wrapping_shl with counts >= 32 ---"); + { + let input = extremes(len); + let want = expected_u32(&input, &past_width, |v, c| v.wrapping_shl(c)); + let input_dev = DeviceBuffer::from_host(&stream, &input).unwrap(); + let count_dev = DeviceBuffer::from_host(&stream, &past_width).unwrap(); + let mut out_dev = DeviceBuffer::from_host(&stream, &invert(&want)).unwrap(); + + // SAFETY: launch shape matches the buffers; every lane stays in range. + unsafe { + module.wrapping_shl_u32(stream.as_ref(), cfg, &input_dev, &count_dev, &mut out_dev) + } + .expect("wrapping_shl_u32 launch failed"); + + let got = out_dev.to_host_vec(&stream).unwrap(); + if got == want { + println!("✓ counts 32/33/63 mask to 0/1/31 for all {len} lanes"); + } else { + println!("✗ wrapping shift mismatch"); + println!(" got {:?}", &got[..8]); + println!(" want {:?}", &want[..8]); + failed = true; + } + } + + if failed { + std::process::exit(1); + } + println!("\nSUCCESS: shifted warp reductions match the host on all lanes"); +}