From 6145cb9706f15bac5402ddaff43daea9ae657003 Mon Sep 17 00:00:00 2001 From: 0z5a Date: Sat, 26 Sep 2026 14:55:49 +0800 Subject: [PATCH 1/2] fix(mir-lower): unify shift operands on the dialect's integer representation `convert_shift` reconciled the shift value and the shift count by comparing widths: widen with `zext` when the value is wider, otherwise narrow with `trunc`. The `trunc` arm is also reached when the widths are *equal* and only the dialect representation differs, and `llvm.trunc i32 -> ui32` is not a legal instruction - LLVM requires a strictly smaller result. #1328 reports the whole device module failing verification for a `redux.sync` result shifted by a literal. Two representation rules decide the fix, and the second is not optional: - Every op in the chain is an `IntBinArithOp`, and pliron-llvm rejects a represented (`uiN`/`siN`) operand of one outright. A represented shift value is therefore unified to the signless representation first, as a same-width change of the same bits. Carrying it into the count mask cannot produce a legal module whatever the count looks like. - A count that differs only in representation then takes that same same-width change rather than a truncation. Handing it over unchanged is not enough either: the mask and the shift are `SameOperandsAndResultType` ops, so the whole chain has to agree on one representation. The reported reproducer already passes on main - #1264 started translating reduction results through `convert_type`, which removed the only path that produced a represented value here. The illegal arm stayed, so this also adds the regression that reaches it. The count normalization is callable on its own, and the tests drive it together with the real count mask and the real `convert_shl` / `convert_shr` from operands built at that boundary, because the source pipeline normalizes both operands before a shift is converted. Signed-off-by: 0z5a --- .../mir-lower/src/convert/ops/arithmetic.rs | 486 +++++++++++++++++- crates/mir-lower/tests/lowering_test/main.rs | 1 + .../lowering_test/shift_representation.rs | 434 ++++++++++++++++ 3 files changed, 898 insertions(+), 23 deletions(-) create mode 100644 crates/mir-lower/tests/lowering_test/shift_representation.rs diff --git a/crates/mir-lower/src/convert/ops/arithmetic.rs b/crates/mir-lower/src/convert/ops/arithmetic.rs index 0c16f6f7e0..c79032061d 100644 --- a/crates/mir-lower/src/convert/ops/arithmetic.rs +++ b/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/crates/mir-lower/tests/lowering_test/main.rs b/crates/mir-lower/tests/lowering_test/main.rs index 09ab58e43a..1e93617602 100644 --- a/crates/mir-lower/tests/lowering_test/main.rs +++ b/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/crates/mir-lower/tests/lowering_test/shift_representation.rs b/crates/mir-lower/tests/lowering_test/shift_representation.rs new file mode 100644 index 0000000000..2923a4aa99 --- /dev/null +++ b/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) + ); +} From 0c56bbc113a270f689e22c9dca0d5c60a685a37c Mon Sep 17 00:00:00 2001 From: 0z5a Date: Sat, 26 Sep 2026 14:55:53 +0800 Subject: [PATCH 2/2] example(redux_shift_regression): shifted warp reductions on sm_80 `redux_minmax` writes a reduction straight to memory and `redux_sum` consumes one without an operator between, so nothing covered the shape #1328 reported: a shift applied to a `redux.sync` result. Four kernels over a full warp, 128 threads so four warps share a block, and each warp holds different values, so a reduction that read a neighbour's lanes could not produce the expected answers. - `redux.sync.max.u32` then `<<` and `>>` over three input patterns: a ramp offset per warp, one carrying 0, 1, 0x8000_0000 and 0xffff_ffff in every warp, and fixed-seed random values; - `redux.sync.min.s32` then `>>`, the control for signedness; - `wrapping_shl` with counts 32, 33 and 63. Every lane writes its own result, because `redux.sync` broadcasts to the lanes in the member mask, and each output buffer starts at the bitwise inverse of the host oracle, so a lane the kernel never writes cannot read back as a correct answer. This is a device smoke test. The count-masking invariant itself is asserted by the converter tests, not here: `count & 31` and `wrapping_shl` in the kernel would agree with the oracle even if the backend dropped the mask. Signed-off-by: 0z5a --- .../redux_shift_regression/Cargo.lock | 724 ++++++++++++++++++ .../redux_shift_regression/Cargo.toml | 24 + .../redux_shift_regression/src/main.rs | 298 +++++++ 3 files changed, 1046 insertions(+) create mode 100644 crates/rustc-codegen-cuda/examples/redux_shift_regression/Cargo.lock create mode 100644 crates/rustc-codegen-cuda/examples/redux_shift_regression/Cargo.toml create mode 100644 crates/rustc-codegen-cuda/examples/redux_shift_regression/src/main.rs diff --git a/crates/rustc-codegen-cuda/examples/redux_shift_regression/Cargo.lock b/crates/rustc-codegen-cuda/examples/redux_shift_regression/Cargo.lock new file mode 100644 index 0000000000..c23ab7db71 --- /dev/null +++ b/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/crates/rustc-codegen-cuda/examples/redux_shift_regression/Cargo.toml b/crates/rustc-codegen-cuda/examples/redux_shift_regression/Cargo.toml new file mode 100644 index 0000000000..992b15da65 --- /dev/null +++ b/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/crates/rustc-codegen-cuda/examples/redux_shift_regression/src/main.rs b/crates/rustc-codegen-cuda/examples/redux_shift_regression/src/main.rs new file mode 100644 index 0000000000..ade625f907 --- /dev/null +++ b/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"); +}