Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -589,6 +589,30 @@ pub(super) fn translate_aggregate_rvalue(
})?;
let tuple_ty = types::translate_type(ctx, &rust_tuple_ty)?;

// Normalize pointer representations to the tuple's declared
// element types before constructing the aggregate.
let expected_types = {
let tuple_ref = tuple_ty.deref(ctx);
tuple_ref
.downcast_ref::<dialect_mir::types::MirTupleType>()
.expect("translated tuple must have MirTupleType")
.get_types()
.to_vec()
};

for (value, expected_type) in element_values.iter_mut().zip(expected_types) {
let (normalized, prev_after_cast) = cast_to_expected_pointer_type_if_needed(
ctx,
*value,
expected_type,
block_ptr,
current_prev_op,
loc.clone(),
);
*value = normalized;
current_prev_op = prev_after_cast;
}

// Create mir.construct_tuple operation
use dialect_mir::ops::MirConstructTupleOp;

Expand Down
42 changes: 42 additions & 0 deletions cuda-oxide/crates/mir-importer/src/translator/rvalue/coerce.rs
Original file line number Diff line number Diff line change
Expand Up @@ -342,6 +342,48 @@ mod tests {
use dialect_mir::types::{MirPointerKind, MirPtrType};
use pliron::builtin::types::{IntegerType, Signedness};

#[test]
fn expected_pointer_normalization_preserves_cluster_shared_raw_const() {
let mut ctx = Context::new();
crate::translator::register_dialects(&mut ctx);

let pointee: TypeHandle = IntegerType::get(&ctx, 32, Signedness::Unsigned).into();

let source_ty: TypeHandle = MirPtrType::get_with_kind(
&mut ctx,
pointee,
false,
dialect_mir::types::address_space::CLUSTER_SHARED,
MirPointerKind::RawConst,
)
.into();

let target_ty: TypeHandle =
MirPtrType::get_generic_with_kind(&mut ctx, pointee, false, MirPointerKind::RawConst)
.into();

let block = BasicBlock::new(&mut ctx, None, vec![source_ty]);
let source_value = block.deref(&ctx).get_argument(0);

let (normalized, last_op) = cast_to_expected_pointer_type_if_needed(
&mut ctx,
source_value,
target_ty,
block,
None,
Location::Unknown,
);

assert_eq!(normalized.get_type(&ctx), target_ty);

let cast_op = last_op.expect("AS7 to generic normalization must insert a pointer cast");

assert!(
Operation::get_op::<MirCastOp>(cast_op, &ctx).is_some(),
"normalization must produce mir.cast"
);
}

#[test]
fn expected_pointer_normalization_rejects_concrete_kind_change() {
let mut ctx = Context::new();
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,7 @@
//! Uses unified compilation: single `cargo oxide run cluster`

use core::ptr::{addr_of, addr_of_mut};
use cuda_device::{DisjointSlice, SharedArray, cluster, cluster_launch, kernel, thread};
use cuda_device::{DisjointSlice, SharedArray, cluster, cluster_launch, device, kernel, thread};
use cuda_host::cuda_module;

// ============================================================================
Expand Down Expand Up @@ -136,6 +136,14 @@ mod kernels {
// Test 3: Distributed Shared Memory (Ring Exchange)
// ============================================================================

#[inline(never)]
#[device]
fn bounce_cluster_pair(
pair: (*const u32, u64),
) -> (*const u32, u64) {
pair
}

/// Test kernel for distributed shared memory via map_shared_rank() + dereference.
#[kernel]
#[cluster_launch(4, 1, 1)]
Expand All @@ -160,7 +168,18 @@ mod kernels {
let neighbor_rank = (my_rank + 1) % cluster_size;
let neighbor_ptr =
unsafe { cluster::map_shared_rank(addr_of!(SHMEM) as *const u32, neighbor_rank) };
let neighbor_value = unsafe { *neighbor_ptr };
let pair =
bounce_cluster_pair((neighbor_ptr, neighbor_rank as u64));

let direct = unsafe { *neighbor_ptr };
let bounced = unsafe { *pair.0 };

let neighbor_value =
if bounced == direct && pair.1 == neighbor_rank as u64 {
bounced
} else {
u32::MAX
};

let idx = my_rank as usize;
if idx < output.len() {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -71,4 +71,9 @@ require_entry_shape test_dsmem_ring_exchange \
require_entry_shape test_dsmem_mapped_store \
"cluster-shared remote store" 'st\.shared::cluster\.'

# Verify the tuple round-trip call before NVVM optimizations.
# The PTX optimizer may eliminate or inline this call.
require_ll_shape "tuple round-trip LLVM call" \
'call \{ ptr, i64 \} @bounce_cluster_pair\('

echo "cluster code shape: PASS"