use crate::cli::Target;
use crate::ir::{CudaQualifier, IrNode, TranslationUnit};
use crate::targets::TargetBackend;
pub fn emit(target: Target, backend: &dyn TargetBackend, unit: &TranslationUnit) -> String {
let source = unit.source.clone();
let banner = backend.banner(&source);
let replaced = apply_replacements(target, unit);
format!("{banner}{replaced}")
}
pub fn apply_replacements(target: Target, unit: &TranslationUnit) -> String {
let mut s = unit.source.clone();
let mut edits: Vec<(usize, usize, String)> = Vec::with_capacity(unit.nodes.len());
for node in &unit.nodes {
if let Some(repl) = replacement_for(target, node) {
let end = effective_end(target, node, &s);
edits.push((node.start(), end, repl));
}
}
edits.sort_by_key(|(start, end, _)| (*start, std::cmp::Reverse(*end)));
let mut filtered: Vec<(usize, usize, String)> = Vec::with_capacity(edits.len());
let mut last_end: usize = 0;
for (start, end, text) in edits {
if start >= last_end {
filtered.push((start, end, text));
last_end = end;
} else {
}
}
edits = filtered;
edits.sort_by_key(|(start, _, _)| *start);
let mut acc_delta: i64 = 0;
for (start, end, text) in edits {
let orig_len = (end - start) as i64;
let adj_start = (start as i64 + acc_delta).max(0) as usize;
let adj_end = (end as i64 + acc_delta).max(0) as usize;
if adj_end > s.len() || adj_start > adj_end {
continue;
}
s.replace_range(adj_start..adj_end, &text);
acc_delta += text.len() as i64 - orig_len;
}
s
}
pub fn replacement_for(target: Target, node: &IrNode) -> Option<String> {
let slot = crate::ir::slot_index(target);
match node {
IrNode::QualifierDecl { qualifier, .. } => qualifier_rewrite(*qualifier, target),
IrNode::KernelLaunch {
kernel,
grid,
block,
smem,
stream,
args,
..
} => Some(launch_rewrite(target, kernel, grid, block, smem, stream, args)),
IrNode::RuntimeCall { name: _, args, mappings, .. } => {
if let Some(mapped) = mappings[slot].clone() {
let args_joined = args.join(", ");
Some(format!("{mapped}({args_joined})"))
} else {
None
}
}
IrNode::BuiltinRef { kind, .. } => Some(builtin_rewrite(*kind, target)),
IrNode::HeaderInclude { header, replacements, .. } => {
replacements[slot].clone().map(|mapped| match target {
Target::Rust => format!("// was: #include {header} -> {mapped}"),
_ => format!("#include {mapped} /* was: {header} */"),
})
}
IrNode::AtomicIntrinsic { name, .. } => Some(atomic_rewrite(name, target)),
IrNode::KernelDef { .. } => None, IrNode::Warning { .. } => None,
}
}
pub fn effective_end(target: Target, node: &IrNode, source: &str) -> usize {
if let IrNode::BuiltinRef {
end,
has_field_access: true,
..
} = node
{
if target != Target::Hip {
let bytes = source.as_bytes();
let mut p = *end;
while p < bytes.len() && bytes[p].is_ascii_whitespace() {
p += 1;
}
if p + 1 < bytes.len()
&& bytes[p] == b'.'
&& matches!(bytes[p + 1], b'x' | b'y' | b'z')
&& (p + 2 == bytes.len() || !(bytes[p + 2].is_ascii_alphanumeric() || bytes[p + 2] == b'_'))
{
return p + 2;
}
}
}
node.end()
}
fn qualifier_rewrite(q: CudaQualifier, target: Target) -> Option<String> {
use CudaQualifier::*;
let s = match (q, target) {
(Global, Target::Hip) => "__global__",
(Global, Target::Sycl) => "// TODO(decuda): rewrite as SYCL kernel lambda\n",
(Global, Target::Rust) => "// TODO(decuda): rewrite as rust-gpu kernel fn\n",
(Global, Target::Opencl) => "__kernel",
(Device, Target::Hip) => "__device__",
(Device, Target::Sycl) => "// SYCL device function",
(Device, Target::Rust) => "/* device function */",
(Device, Target::Opencl) => "__device",
(Host, Target::Hip) => "__host__",
(Host, Target::Sycl) => "",
(Host, Target::Rust) => "",
(Host, Target::Opencl) => "",
(ForceInline, Target::Hip) => "__forceinline__",
(ForceInline, Target::Sycl) => "[[clang::always_inline]]",
(ForceInline, Target::Rust) => "#[inline(always)]",
(ForceInline, Target::Opencl) => "",
(NoInline, Target::Hip) => "__noinline__",
(NoInline, Target::Sycl) => "[[gnu::noinline]]",
(NoInline, Target::Rust) => "#[inline(never)]",
(NoInline, Target::Opencl) => "",
(Shared, Target::Hip) => "__shared__",
(Shared, Target::Sycl) => "__shared__", (Shared, Target::Rust) => "/* shared -> rust-gpu group_memory */",
(Shared, Target::Opencl) => "__local",
(Constant, Target::Hip) => "__constant__",
(Constant, Target::Sycl) => "/* constant -> SYCL constant_accessor */",
(Constant, Target::Rust) => "/* constant */",
(Constant, Target::Opencl) => "__constant",
(Managed, Target::Hip) => "__managed__",
(Managed, Target::Sycl) => "/* managed -> SYCL USM */",
(Managed, Target::Rust) => "/* managed -> unified memory */",
(Managed, Target::Opencl) => "/* use SVM: clSVMAlloc */",
(Restricted, Target::Hip) => "__restrict__",
(Restricted, Target::Sycl) => "",
(Restricted, Target::Rust) => "",
(Restricted, Target::Opencl) => "restrict",
(Texture, _) => return None,
(Surface, _) => return None,
(LaunchBounds, _) => return None,
(ClusterDim, _) => return None,
(GridDim, _) => return None,
(Const, _) => "",
(Pinned, _) => "",
(_, Target::All) => "",
};
Some(s.to_string())
}
fn launch_rewrite(
target: Target,
kernel: &str,
grid: &str,
block: &str,
smem: &Option<String>,
stream: &Option<String>,
args: &[String],
) -> String {
let args_joined = args.join(", ");
match target {
Target::Hip => {
let smem = smem.as_deref().unwrap_or("0");
let stream = stream.as_deref().unwrap_or("0");
format!(
"hipLaunchKernelGGL({kernel}, dim3({grid}), dim3({block}), {smem}, {stream}, {args_joined})"
)
}
Target::Sycl => {
let smem = smem.as_deref().unwrap_or("none");
let stream = stream.as_deref().unwrap_or("default");
format!(
"{{ /* decuda SYCL launch: queue.submit([&](sycl::handler& h) {{ h.parallel_for(sycl::range<3>{{{grid}}}, [=](sycl::item<3> it) {{ /* kernel `{kernel}` body with thread indices from it.get_*() */ }}); }}); smem={smem} stream={stream} args={args_joined} */ }}"
)
}
Target::Rust => {
let smem = smem.as_deref().unwrap_or("0");
let stream_v = stream.as_deref().unwrap_or("default");
format!(
"{{ /* decuda cust launch */ let _kernel = modules.get_function(\"{kernel}\"); unsafe {{ let _ = launch!( _kernel<<<{grid} as grid_size, {block} as block_size, {smem} as usize, {stream_v}>>>({args_joined}) ); }} }}"
)
}
Target::Opencl => {
format!(
"clEnqueueNDRangeKernel(queue, {kernel}_kernel, 1, NULL, (size_t[1]){{{grid}}}, (size_t[1]){{{block}}}, 0, NULL, NULL) /* args: {args_joined} */",
)
}
Target::All => format!("{kernel}<<<{grid}, {block}>>>({args_joined})"),
}
}
fn builtin_rewrite(kind: crate::ir::BuiltinKind, target: Target) -> String {
use crate::ir::BuiltinKind::*;
let s = match (kind, target) {
(ThreadIdx, Target::Hip) => "threadIdx",
(ThreadIdx, Target::Sycl) => "item.get_local_id()",
(ThreadIdx, Target::Rust) => "thread_idx",
(ThreadIdx, Target::Opencl) => "get_local_id(0)",
(BlockIdx, Target::Hip) => "blockIdx",
(BlockIdx, Target::Sycl) => "item.get_group(0)",
(BlockIdx, Target::Rust) => "block_idx",
(BlockIdx, Target::Opencl) => "get_group_id(0)",
(BlockDim, Target::Hip) => "blockDim",
(BlockDim, Target::Sycl) => "item.get_local_range()",
(BlockDim, Target::Rust) => "block_dim",
(BlockDim, Target::Opencl) => "get_local_size(0)",
(GridDim, Target::Hip) => "gridDim",
(GridDim, Target::Sycl) => "item.get_global_range() / item.get_local_range()",
(GridDim, Target::Rust) => "grid_dim",
(GridDim, Target::Opencl) => "get_num_groups(0)",
(WarpSize, Target::Hip) => "warpSize",
(WarpSize, Target::Sycl) => "32 /* warpSize */",
(WarpSize, Target::Rust) => "WARP_SIZE",
(WarpSize, Target::Opencl) => "warpSize /* CL_DEVICE_WARP_SIZE_NV */",
(SyncThreads, Target::Hip) => "__syncthreads()",
(SyncThreads, Target::Sycl) => "item.barrier(sycl::access::fence_space::global_space)",
(SyncThreads, Target::Rust) => "group.sync()",
(SyncThreads, Target::Opencl) => "barrier(CLK_LOCAL_MEM_FENCE)",
(SyncWarp, Target::Hip) => "__syncwarp(0xFFFFFFFFu)",
(SyncWarp, Target::Sycl) => "/* no native syncwarp on SYCL */ item.barrier() /* fallback */",
(SyncWarp, Target::Rust) => "/* syncwarp: only lane=0 of warp at once */",
(SyncWarp, Target::Opencl) => "barrier(CLK_LOCAL_MEM_FENCE) /* approx */",
(FsyncBlock, Target::Hip) => "__sync_block()",
(FsyncBlock, Target::Sycl) => "item.barrier()",
(FsyncBlock, Target::Rust) => "group.sync()",
(FsyncBlock, Target::Opencl) => "barrier(CLK_LOCAL_MEM_FENCE)",
(LaneId, Target::Hip) => "__laneid()",
(LaneId, Target::Sycl) => "item.get_sub_group().get_local_id()",
(LaneId, Target::Rust) => "lane_id",
(LaneId, Target::Opencl) => "get_sub_group_id() * get_sub_group_size() + get_sub_group_local_id()",
(WarpId, Target::Hip) => "__warp_id()",
(WarpId, Target::Sycl) => "item.get_sub_group().get_group_id()",
(WarpId, Target::Rust) => "warp_id",
(WarpId, Target::Opencl) => "get_sub_group_id()",
(ShflSync, Target::Hip) => "__shfl_sync",
(ShflSync, Target::Sycl) => "sycl::sub_group::shuffle",
(ShflSync, Target::Rust) => "/* TODO: __shfl_sync */",
(ShflSync, Target::Opencl) => "sub_group_shuffle",
(BallotSync, Target::Hip) => "__ballot_sync",
(BallotSync, Target::Sycl) => "sycl::sub_group::ballot",
(BallotSync, Target::Rust) => "/* TODO: __ballot_sync */",
(BallotSync, Target::Opencl) => "sub_group_ballot",
(AnySync, Target::Hip) => "__any_sync",
(AnySync, Target::Sycl) => "sycl::sub_group::any",
(AnySync, Target::Rust) => "/* TODO: __any_sync */",
(AnySync, Target::Opencl) => "sub_group_any",
(AllSync, Target::Hip) => "__all_sync",
(AllSync, Target::Sycl) => "sycl::sub_group::all",
(AllSync, Target::Rust) => "/* TODO: __all_sync */",
(AllSync, Target::Opencl) => "sub_group_all",
(ActiveMask, Target::Hip) => "__activemask",
(ActiveMask, Target::Sycl) => "sycl::sub_group::get_local_range",
(ActiveMask, Target::Rust) => "/* TODO: __activemask */",
(ActiveMask, Target::Opencl) => "get_sub_group_size",
(_, Target::All) => "/* decuda: unsupported builtin */",
};
s.to_string()
}
fn atomic_rewrite(name: &str, target: Target) -> String {
match target {
Target::Hip | Target::Opencl => name.to_string(),
Target::Sycl | Target::Rust => name.to_string(),
Target::All => name.to_string(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ir::{BuiltinKind, CudaQualifier};
#[test]
fn global_qualifier_to_opencl_kernel() {
let s = qualifier_rewrite(CudaQualifier::Global, Target::Opencl).unwrap();
assert_eq!(s, "__kernel");
}
#[test]
fn threadidx_rewrite_per_target() {
assert_eq!(builtin_rewrite(BuiltinKind::ThreadIdx, Target::Hip), "threadIdx");
assert_eq!(
builtin_rewrite(BuiltinKind::ThreadIdx, Target::Opencl),
"get_local_id(0)"
);
}
#[test]
fn syncthreads_rewrites_globally() {
assert!(builtin_rewrite(BuiltinKind::SyncThreads, Target::Sycl).contains("barrier"));
}
#[test]
fn shfl_sync_per_target() {
assert_eq!(builtin_rewrite(BuiltinKind::ShflSync, Target::Hip), "__shfl_sync");
assert_eq!(builtin_rewrite(BuiltinKind::ShflSync, Target::Opencl), "sub_group_shuffle");
assert!(builtin_rewrite(BuiltinKind::ShflSync, Target::Sycl).contains("shuffle"));
assert!(builtin_rewrite(BuiltinKind::ShflSync, Target::Rust).contains("TODO"));
}
#[test]
fn ballot_sync_per_target() {
assert_eq!(builtin_rewrite(BuiltinKind::BallotSync, Target::Hip), "__ballot_sync");
assert_eq!(builtin_rewrite(BuiltinKind::BallotSync, Target::Opencl), "sub_group_ballot");
}
#[test]
fn activemask_per_target() {
assert_eq!(builtin_rewrite(BuiltinKind::ActiveMask, Target::Hip), "__activemask");
assert_eq!(builtin_rewrite(BuiltinKind::ActiveMask, Target::Opencl), "get_sub_group_size");
}
}