Skip to main content

rustc_codegen_llvm/builder/
gpu_offload.rs

1use std::ffi::CString;
2
3use bitflags::Flags;
4use llvm::Linkage::*;
5use rustc_abi::Align;
6use rustc_codegen_ssa::MemFlags;
7use rustc_codegen_ssa::common::TypeKind;
8use rustc_codegen_ssa::mir::operand::{OperandRef, OperandValue};
9use rustc_codegen_ssa::traits::{BaseTypeCodegenMethods, BuilderMethods, ReturnSlot};
10use rustc_middle::ty::offload_meta::{MappingFlags, OffloadMetadata, OffloadSize};
11use rustc_span::bug;
12
13use crate::attributes;
14use crate::builder::Builder;
15use crate::common::CodegenCx;
16use crate::context::SimpleCx;
17use crate::llvm::AttributePlace::Function;
18use crate::llvm::{self, Linkage, Type, Value};
19
20// LLVM kernel-independent globals required for offloading
21pub(crate) struct OffloadGlobals<'ll> {
22    pub launcher_fn: &'ll llvm::Value,
23    pub launcher_ty: &'ll llvm::Type,
24
25    pub kernel_args_ty: &'ll llvm::Type,
26
27    pub offload_entry_ty: &'ll llvm::Type,
28
29    pub begin_mapper: &'ll llvm::Value,
30    pub end_mapper: &'ll llvm::Value,
31    pub mapper_fn_ty: &'ll llvm::Type,
32
33    pub ident_t_global: &'ll llvm::Value,
34}
35
36impl<'ll> OffloadGlobals<'ll> {
37    pub(crate) fn declare(cx: &CodegenCx<'ll, '_>) -> Self {
38        let (launcher_fn, launcher_ty) = generate_launcher(cx);
39        let kernel_args_ty = KernelArgsTy::new_decl(cx);
40        let offload_entry_ty = TgtOffloadEntry::new_decl(cx);
41        let (begin_mapper, _, end_mapper, mapper_fn_ty) = gen_tgt_data_mappers(cx);
42        let ident_t_global = generate_at_one(cx);
43
44        // We want LLVM's openmp-opt pass to pick up and optimize this module, since it covers both
45        // openmp and offload optimizations.
46        llvm::add_module_flag_u32(cx.llmod(), llvm::ModuleFlagMergeBehavior::Max, "openmp", 51);
47
48        OffloadGlobals {
49            launcher_fn,
50            launcher_ty,
51            kernel_args_ty,
52            offload_entry_ty,
53            begin_mapper,
54            end_mapper,
55            mapper_fn_ty,
56            ident_t_global,
57        }
58    }
59}
60
61pub(crate) struct OffloadKernelDims<'ll> {
62    num_workgroups: &'ll Value,
63    threads_per_block: &'ll Value,
64    workgroup_dims: &'ll Value,
65    thread_dims: &'ll Value,
66}
67
68impl<'ll> OffloadKernelDims<'ll> {
69    pub(crate) fn from_operands<'tcx>(
70        builder: &mut Builder<'_, 'll, 'tcx>,
71        workgroup_op: &OperandRef<'tcx, &'ll llvm::Value>,
72        thread_op: &OperandRef<'tcx, &'ll llvm::Value>,
73    ) -> Self {
74        let cx = builder.cx;
75        let arr_ty = cx.type_array(cx.type_i32(), 3);
76        let four = Align::from_bytes(4).unwrap();
77
78        let OperandValue::Ref(place) = workgroup_op.val else {
79            ::rustc_span::macros::bug_impl(None,
    format_args!("expected array operand by reference"), Location::caller());bug!("expected array operand by reference");
80        };
81        let workgroup_val = builder.load(arr_ty, place.llval, four);
82
83        let OperandValue::Ref(place) = thread_op.val else {
84            ::rustc_span::macros::bug_impl(None,
    format_args!("expected array operand by reference"), Location::caller());bug!("expected array operand by reference");
85        };
86        let thread_val = builder.load(arr_ty, place.llval, four);
87
88        fn mul_dim3<'ll, 'tcx>(
89            builder: &mut Builder<'_, 'll, 'tcx>,
90            arr: &'ll Value,
91        ) -> &'ll Value {
92            let x = builder.extract_value(arr, 0);
93            let y = builder.extract_value(arr, 1);
94            let z = builder.extract_value(arr, 2);
95
96            let xy = builder.mul(x, y);
97            builder.mul(xy, z)
98        }
99
100        let num_workgroups = mul_dim3(builder, workgroup_val);
101        let threads_per_block = mul_dim3(builder, thread_val);
102
103        OffloadKernelDims {
104            workgroup_dims: workgroup_val,
105            thread_dims: thread_val,
106            num_workgroups,
107            threads_per_block,
108        }
109    }
110}
111
112// ; Function Attrs: nounwind
113// declare i32 @__tgt_target_kernel(ptr, i64, i32, i32, ptr, ptr) #2
114fn generate_launcher<'ll>(cx: &CodegenCx<'ll, '_>) -> (&'ll llvm::Value, &'ll llvm::Type) {
115    let tptr = cx.type_ptr();
116    let ti64 = cx.type_i64();
117    let ti32 = cx.type_i32();
118    let args = ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
        [tptr, ti64, ti32, ti32, tptr, tptr]))vec![tptr, ti64, ti32, ti32, tptr, tptr];
119    let tgt_fn_ty = cx.type_func(&args, ti32);
120    let name = "__tgt_target_kernel";
121    let tgt_decl = declare_offload_fn(&cx, name, tgt_fn_ty);
122    let nounwind = llvm::AttributeKind::NoUnwind.create_attr(cx.llcx);
123    attributes::apply_to_llfn(tgt_decl, Function, &[nounwind]);
124    (tgt_decl, tgt_fn_ty)
125}
126
127/// Declares the `omp_get_num_devices` runtime function and returns the
128/// declaration together with its type.
129pub(crate) fn declare_omp_get_num_devices<'ll>(
130    cx: &CodegenCx<'ll, '_>,
131) -> (&'ll llvm::Value, &'ll llvm::Type) {
132    let ti32 = cx.type_i32();
133    let tgt_fn_ty = cx.type_func(&[], ti32);
134    let name = "omp_get_num_devices";
135    let tgt_decl = declare_offload_fn(&cx, name, tgt_fn_ty);
136    let nounwind = llvm::AttributeKind::NoUnwind.create_attr(cx.llcx);
137    attributes::apply_to_llfn(tgt_decl, Function, &[nounwind]);
138    (tgt_decl, tgt_fn_ty)
139}
140
141// What is our @1 here? A magic global, used in our data_{begin/update/end}_mapper:
142// @0 = private unnamed_addr constant [23 x i8] c";unknown;unknown;0;0;;\00", align 1
143// @1 = private unnamed_addr constant %struct.ident_t { i32 0, i32 2, i32 0, i32 22, ptr @0 }, align 8
144// FIXME(offload): @0 should include the file name (e.g. lib.rs) in which the function to be
145// offloaded was defined.
146pub(crate) fn generate_at_one<'ll>(cx: &CodegenCx<'ll, '_>) -> &'ll llvm::Value {
147    let unknown_txt = ";unknown;unknown;0;0;;";
148    let c_entry_name = CString::new(unknown_txt).unwrap();
149    let c_val = c_entry_name.as_bytes_with_nul();
150    let initializer = crate::common::bytes_in_context(cx.llcx, c_val);
151    let at_zero = add_unnamed_global(&cx, &"", initializer, PrivateLinkage);
152    llvm::set_alignment(at_zero, Align::ONE);
153
154    // @1 = private unnamed_addr constant %struct.ident_t { i32 0, i32 2, i32 0, i32 22, ptr @0 }, align 8
155    let struct_ident_ty = cx.type_named_struct("struct.ident_t");
156    let struct_elems = ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
        [cx.get_const_i32(0), cx.get_const_i32(2), cx.get_const_i32(0),
                cx.get_const_i32(22), at_zero]))vec![
157        cx.get_const_i32(0),
158        cx.get_const_i32(2),
159        cx.get_const_i32(0),
160        cx.get_const_i32(22),
161        at_zero,
162    ];
163    let struct_elems_ty: Vec<_> = struct_elems.iter().map(|&x| cx.val_ty(x)).collect();
164    let initializer = crate::common::named_struct(struct_ident_ty, &struct_elems);
165    cx.set_struct_body(struct_ident_ty, &struct_elems_ty, false);
166    let at_one = add_unnamed_global(&cx, &"", initializer, PrivateLinkage);
167    llvm::set_alignment(at_one, Align::EIGHT);
168    at_one
169}
170
171pub(crate) struct TgtOffloadEntry {
172    //   uint64_t Reserved;
173    //   uint16_t Version;
174    //   uint16_t Kind;
175    //   uint32_t Flags; Flags associated with the entry (see Target Region Entry Flags)
176    //   void *Address; Address of global symbol within device image (function or global)
177    //   char *SymbolName;
178    //   uint64_t Size; Size of the entry info (0 if it is a function)
179    //   uint64_t Data;
180    //   void *AuxAddr;
181}
182
183impl TgtOffloadEntry {
184    pub(crate) fn new_decl<'ll>(cx: &CodegenCx<'ll, '_>) -> &'ll llvm::Type {
185        let offload_entry_ty = cx.type_named_struct("struct.__tgt_offload_entry");
186        let tptr = cx.type_ptr();
187        let ti64 = cx.type_i64();
188        let ti32 = cx.type_i32();
189        let ti16 = cx.type_i16();
190        // For each kernel to run on the gpu, we will later generate one entry of this type.
191        // copied from LLVM
192        let entry_elements = ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
        [ti64, ti16, ti16, ti32, tptr, tptr, ti64, ti64, tptr]))vec![ti64, ti16, ti16, ti32, tptr, tptr, ti64, ti64, tptr];
193        cx.set_struct_body(offload_entry_ty, &entry_elements, false);
194        offload_entry_ty
195    }
196
197    fn new<'ll>(
198        cx: &CodegenCx<'ll, '_>,
199        region_id: &'ll Value,
200        llglobal: &'ll Value,
201    ) -> [&'ll Value; 9] {
202        let reserved = cx.get_const_i64(0);
203        let version = cx.get_const_i16(1);
204        let kind = cx.get_const_i16(1);
205        let flags = cx.get_const_i32(0);
206        let size = cx.get_const_i64(0);
207        let data = cx.get_const_i64(0);
208        let aux_addr = cx.const_null(cx.type_ptr());
209        [reserved, version, kind, flags, region_id, llglobal, size, data, aux_addr]
210    }
211}
212
213// Taken from the LLVM APITypes.h declaration:
214struct KernelArgsTy {
215    //  uint32_t Version = 0; // Version of this struct for ABI compatibility.
216    //  uint32_t NumArgs = 0; // Number of arguments in each input pointer.
217    //  void **ArgBasePtrs =
218    //      nullptr;                 // Base pointer of each argument (e.g. a struct).
219    //  void **ArgPtrs = nullptr;    // Pointer to the argument data.
220    //  int64_t *ArgSizes = nullptr; // Size of the argument data in bytes.
221    //  int64_t *ArgTypes = nullptr; // Type of the data (e.g. to / from).
222    //  void **ArgNames = nullptr;   // Name of the data for debugging, possibly null.
223    //  void **ArgMappers = nullptr; // User-defined mappers, possibly null.
224    //  uint64_t Tripcount =
225    // 0; // Tripcount for the teams / distribute loop, 0 otherwise.
226    // struct {
227    //    uint64_t NoWait : 1; // Was this kernel spawned with a `nowait` clause.
228    //    uint64_t IsCUDA : 1; // Was this kernel spawned via CUDA.
229    //    uint64_t Unused : 62;
230    //  } Flags = {0, 0, 0}; // totals to 64 Bit, 8 Byte
231    //  // The number of teams (for x,y,z dimension).
232    //  uint32_t NumTeams[3] = {0, 0, 0};
233    //  // The number of threads (for x,y,z dimension).
234    //  uint32_t ThreadLimit[3] = {0, 0, 0};
235    //  uint32_t DynCGroupMem = 0; // Amount of dynamic cgroup memory requested.
236}
237
238impl KernelArgsTy {
239    const OFFLOAD_VERSION: u64 = 3;
240    const FLAGS: u64 = 1 << 6; // Enable StrictBlocksAndThreads
241    const TRIPCOUNT: u64 = 0;
242    fn new_decl<'ll>(cx: &CodegenCx<'ll, '_>) -> &'ll Type {
243        let kernel_arguments_ty = cx.type_named_struct("struct.__tgt_kernel_arguments");
244        let tptr = cx.type_ptr();
245        let ti64 = cx.type_i64();
246        let ti32 = cx.type_i32();
247        let tarr = cx.type_array(ti32, 3);
248
249        let kernel_elements =
250            ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
        [ti32, ti32, tptr, tptr, tptr, tptr, tptr, tptr, ti64, ti64, tarr,
                tarr, ti32]))vec![ti32, ti32, tptr, tptr, tptr, tptr, tptr, tptr, ti64, ti64, tarr, tarr, ti32];
251
252        cx.set_struct_body(kernel_arguments_ty, &kernel_elements, false);
253        kernel_arguments_ty
254    }
255
256    fn new<'ll, 'tcx>(
257        cx: &CodegenCx<'ll, 'tcx>,
258        num_args: u64,
259        memtransfer_types: &'ll Value,
260        geps: [&'ll Value; 3],
261        workgroup_dims: &'ll Value,
262        thread_dims: &'ll Value,
263        dyn_cache: &'ll Value,
264    ) -> [(Align, &'ll str, &'ll Value); 13] {
265        let four = Align::from_bytes(4).expect("4 Byte alignment should work");
266        let eight = Align::EIGHT;
267
268        [
269            (four, "Version", cx.get_const_i32(KernelArgsTy::OFFLOAD_VERSION)),
270            (four, "NumArgs", cx.get_const_i32(num_args)),
271            (eight, "ArgBasePtrs", geps[0]),
272            (eight, "ArgPtrs", geps[1]),
273            (eight, "ArgSizes", geps[2]),
274            (eight, "ArgTypes", memtransfer_types),
275            // The next two are debug infos. FIXME(offload): set them
276            (eight, "ArgNames", cx.const_null(cx.type_ptr())), // dbg
277            (eight, "ArgMappers", cx.const_null(cx.type_ptr())), // dbg
278            (eight, "Tripcount", cx.get_const_i64(KernelArgsTy::TRIPCOUNT)),
279            (eight, "Flags", cx.get_const_i64(KernelArgsTy::FLAGS)),
280            (four, "NumTeams", workgroup_dims),
281            (four, "ThreadLimit", thread_dims),
282            (four, "DynCGroupMem", dyn_cache),
283        ]
284    }
285}
286
287// Contains LLVM values needed to manage offloading for a single kernel.
288#[derive(#[automatically_derived]
impl<'ll> ::core::marker::Copy for OffloadKernelGlobals<'ll> { }Copy, #[automatically_derived]
#[doc(hidden)]
unsafe impl<'ll> ::core::clone::TrivialClone for OffloadKernelGlobals<'ll> { }
#[automatically_derived]
impl<'ll> ::core::clone::Clone for OffloadKernelGlobals<'ll> {
    #[inline]
    fn clone(&self) -> Self {
        let _: ::core::clone::AssertParamIsClone<&'ll llvm::Value>;
        let _: ::core::clone::AssertParamIsClone<&'ll llvm::Value>;
        let _: ::core::clone::AssertParamIsClone<&'ll llvm::Value>;
        let _: ::core::clone::AssertParamIsClone<&'ll llvm::Value>;
        let _: ::core::clone::AssertParamIsClone<&'ll llvm::Value>;
        *self
    }
}Clone)]
289pub(crate) struct OffloadKernelGlobals<'ll> {
290    pub offload_sizes: &'ll llvm::Value,
291    pub memtransfer_begin: &'ll llvm::Value,
292    pub memtransfer_kernel: &'ll llvm::Value,
293    pub memtransfer_end: &'ll llvm::Value,
294    pub region_id: &'ll llvm::Value,
295}
296
297fn gen_tgt_data_mappers<'ll>(
298    cx: &CodegenCx<'ll, '_>,
299) -> (&'ll llvm::Value, &'ll llvm::Value, &'ll llvm::Value, &'ll llvm::Type) {
300    let tptr = cx.type_ptr();
301    let ti64 = cx.type_i64();
302    let ti32 = cx.type_i32();
303
304    let args = ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
        [tptr, ti64, ti32, tptr, tptr, tptr, tptr, tptr, tptr]))vec![tptr, ti64, ti32, tptr, tptr, tptr, tptr, tptr, tptr];
305    let mapper_fn_ty = cx.type_func(&args, cx.type_void());
306    let mapper_begin = "__tgt_target_data_begin_mapper";
307    let mapper_update = "__tgt_target_data_update_mapper";
308    let mapper_end = "__tgt_target_data_end_mapper";
309    let begin_mapper_decl = declare_offload_fn(&cx, mapper_begin, mapper_fn_ty);
310    let update_mapper_decl = declare_offload_fn(&cx, mapper_update, mapper_fn_ty);
311    let end_mapper_decl = declare_offload_fn(&cx, mapper_end, mapper_fn_ty);
312
313    let nounwind = llvm::AttributeKind::NoUnwind.create_attr(cx.llcx);
314    attributes::apply_to_llfn(begin_mapper_decl, Function, &[nounwind]);
315    attributes::apply_to_llfn(update_mapper_decl, Function, &[nounwind]);
316    attributes::apply_to_llfn(end_mapper_decl, Function, &[nounwind]);
317
318    (begin_mapper_decl, update_mapper_decl, end_mapper_decl, mapper_fn_ty)
319}
320
321fn add_priv_unnamed_arr<'ll>(cx: &SimpleCx<'ll>, name: &str, vals: &[u64]) -> &'ll llvm::Value {
322    let ti64 = cx.type_i64();
323    let mut size_val = Vec::with_capacity(vals.len());
324    for &val in vals {
325        size_val.push(cx.get_const_i64(val));
326    }
327    let initializer = cx.const_array(ti64, &size_val);
328    add_unnamed_global(cx, name, initializer, PrivateLinkage)
329}
330
331pub(crate) fn add_unnamed_global<'ll>(
332    cx: &SimpleCx<'ll>,
333    name: &str,
334    initializer: &'ll llvm::Value,
335    l: Linkage,
336) -> &'ll llvm::Value {
337    let llglobal = add_global(cx, name, initializer, l);
338    llvm::LLVMSetUnnamedAddress(llglobal, llvm::UnnamedAddr::Global);
339    llglobal
340}
341
342pub(crate) fn add_global<'ll>(
343    cx: &SimpleCx<'ll>,
344    name: &str,
345    initializer: &'ll llvm::Value,
346    l: Linkage,
347) -> &'ll llvm::Value {
348    let c_name = CString::new(name).unwrap();
349    let llglobal: &'ll llvm::Value = llvm::add_global(cx.llmod, cx.val_ty(initializer), &c_name);
350    llvm::set_global_constant(llglobal, true);
351    llvm::set_linkage(llglobal, l);
352    llvm::set_initializer(llglobal, initializer);
353    llglobal
354}
355
356// This function returns a memtransfer value which encodes how arguments to this kernel shall be
357// mapped to/from the gpu. It also returns a region_id with the name of this kernel, to be
358// concatenated into the list of region_ids.
359pub(crate) fn gen_define_handling<'ll>(
360    cx: &CodegenCx<'ll, '_>,
361    metadata: &[OffloadMetadata],
362    symbol: String,
363    offload_globals: &OffloadGlobals<'ll>,
364) -> OffloadKernelGlobals<'ll> {
365    if let Some(entry) = cx.offload_kernel_cache.borrow().get(&symbol) {
366        return *entry;
367    }
368
369    let offload_entry_ty = offload_globals.offload_entry_ty;
370
371    let (sizes, transfer): (Vec<_>, Vec<_>) =
372        metadata.iter().map(|m| (m.payload_size, m.mode)).unzip();
373    // Our begin mapper should only see simplified information about which args have to be
374    // transferred to the device, the end mapper only about which args should be transferred back.
375    // Any information beyond that makes it harder for LLVM's opt pass to evaluate whether it can
376    // safely move (=optimize) the LLVM-IR location of this data transfer. Only the mapping types
377    // mentioned below are handled, so make sure that we don't generate any other ones.
378    let handled_mappings = MappingFlags::TO
379        | MappingFlags::FROM
380        | MappingFlags::TARGET_PARAM
381        | MappingFlags::LITERAL
382        | MappingFlags::IMPLICIT;
383    for arg in &transfer {
384        if true {
    if !!arg.contains_unknown_bits() {
        ::core::panicking::panic("assertion failed: !arg.contains_unknown_bits()")
    };
};debug_assert!(!arg.contains_unknown_bits());
385        if true {
    if !handled_mappings.contains(*arg) {
        ::core::panicking::panic("assertion failed: handled_mappings.contains(*arg)")
    };
};debug_assert!(handled_mappings.contains(*arg));
386    }
387
388    let valid_begin_mappings = MappingFlags::TO | MappingFlags::LITERAL | MappingFlags::IMPLICIT;
389    let transfer_to: Vec<u64> =
390        transfer.iter().map(|m| m.intersection(valid_begin_mappings).bits()).collect();
391    let transfer_from: Vec<u64> =
392        transfer.iter().map(|m| m.intersection(MappingFlags::FROM).bits()).collect();
393    let valid_kernel_mappings = MappingFlags::LITERAL | MappingFlags::IMPLICIT;
394    // FIXME(offload): add `OMP_MAP_TARGET_PARAM = 0x20` only if necessary
395    let transfer_kernel: Vec<u64> = transfer
396        .iter()
397        .map(|m| (m.intersection(valid_kernel_mappings) | MappingFlags::TARGET_PARAM).bits())
398        .collect();
399
400    let actual_sizes = sizes
401        .iter()
402        .map(|s| match s {
403            OffloadSize::Static(sz) => *sz,
404            // NOTE(Sa4dUs): set `.offload_sizes` entry to 0 for sizes that we determine at runtime, just like clang
405            _ => 0,
406        })
407        .collect::<Vec<_>>();
408    let offload_sizes =
409        add_priv_unnamed_arr(&cx, &::alloc::__export::must_use({
        ::alloc::fmt::format(format_args!(".offload_sizes.{0}", symbol))
    })format!(".offload_sizes.{symbol}"), &actual_sizes);
410    let memtransfer_begin =
411        add_priv_unnamed_arr(&cx, &::alloc::__export::must_use({
        ::alloc::fmt::format(format_args!(".offload_maptypes.{0}.begin",
                symbol))
    })format!(".offload_maptypes.{symbol}.begin"), &transfer_to);
412    let memtransfer_kernel =
413        add_priv_unnamed_arr(&cx, &::alloc::__export::must_use({
        ::alloc::fmt::format(format_args!(".offload_maptypes.{0}.kernel",
                symbol))
    })format!(".offload_maptypes.{symbol}.kernel"), &transfer_kernel);
414    let memtransfer_end =
415        add_priv_unnamed_arr(&cx, &::alloc::__export::must_use({
        ::alloc::fmt::format(format_args!(".offload_maptypes.{0}.end",
                symbol))
    })format!(".offload_maptypes.{symbol}.end"), &transfer_from);
416
417    // Next: For each function, generate these three entries. A weak constant,
418    // the llvm.rodata entry name, and  the llvm_offload_entries value
419
420    let name = ::alloc::__export::must_use({
        ::alloc::fmt::format(format_args!(".{0}.region_id", symbol))
    })format!(".{symbol}.region_id");
421    let initializer = cx.get_const_i8(0);
422    let region_id = add_global(&cx, &name, initializer, WeakAnyLinkage);
423
424    let c_entry_name = CString::new(symbol.clone()).unwrap();
425    let c_val = c_entry_name.as_bytes_with_nul();
426    let offload_entry_name = ::alloc::__export::must_use({
        ::alloc::fmt::format(format_args!(".offloading.entry_name.{0}",
                symbol))
    })format!(".offloading.entry_name.{symbol}");
427
428    let initializer = crate::common::bytes_in_context(cx.llcx, c_val);
429    let llglobal = add_unnamed_global(&cx, &offload_entry_name, initializer, InternalLinkage);
430    llvm::set_alignment(llglobal, Align::ONE);
431    llvm::set_section(llglobal, c".llvm.rodata.offloading");
432
433    let name = ::alloc::__export::must_use({
        ::alloc::fmt::format(format_args!(".offloading.entry.{0}", symbol))
    })format!(".offloading.entry.{symbol}");
434
435    // See the __tgt_offload_entry documentation above.
436    let elems = TgtOffloadEntry::new(&cx, region_id, llglobal);
437
438    let initializer = crate::common::named_struct(offload_entry_ty, &elems);
439    let c_name = CString::new(name).unwrap();
440    let offload_entry = llvm::add_global(cx.llmod, offload_entry_ty, &c_name);
441    llvm::set_global_constant(offload_entry, true);
442    llvm::set_linkage(offload_entry, WeakAnyLinkage);
443    llvm::set_initializer(offload_entry, initializer);
444    llvm::set_alignment(offload_entry, Align::EIGHT);
445    let c_section_name = CString::new("llvm_offload_entries").unwrap();
446    llvm::set_section(offload_entry, &c_section_name);
447
448    cx.add_compiler_used_global(offload_entry);
449
450    let result = OffloadKernelGlobals {
451        offload_sizes,
452        memtransfer_begin,
453        memtransfer_kernel,
454        memtransfer_end,
455        region_id,
456    };
457
458    cx.offload_kernel_cache.borrow_mut().insert(symbol, result);
459
460    result
461}
462
463fn declare_offload_fn<'ll>(
464    cx: &CodegenCx<'ll, '_>,
465    name: &str,
466    ty: &'ll llvm::Type,
467) -> &'ll llvm::Value {
468    crate::declare::declare_simple_fn(
469        cx,
470        name,
471        llvm::CallConv::CCallConv,
472        llvm::UnnamedAddr::No,
473        llvm::Visibility::Default,
474        ty,
475    )
476}
477
478pub(crate) fn scalar_width<'ll>(cx: &'ll SimpleCx<'_>, ty: &'ll Type) -> u64 {
479    match cx.type_kind(ty) {
480        TypeKind::Half
481        | TypeKind::Float
482        | TypeKind::Double
483        | TypeKind::X86_FP80
484        | TypeKind::FP128
485        | TypeKind::PPC_FP128 => cx.float_width(ty) as u64,
486        TypeKind::Integer => cx.int_width(ty),
487        other => ::rustc_span::macros::bug_impl(None,
    format_args!("scalar_width was called on a non scalar type {0:?}", other),
    Location::caller())bug!("scalar_width was called on a non scalar type {other:?}"),
488    }
489}
490
491fn get_runtime_size<'ll, 'tcx>(
492    builder: &mut Builder<'_, 'll, 'tcx>,
493    args: &[&'ll Value],
494    index: usize,
495    meta: &OffloadMetadata,
496) -> &'ll Value {
497    match meta.payload_size {
498        OffloadSize::Slice { element_size } => {
499            let length_idx = index + 1;
500            let length = args[length_idx];
501            let length_i64 = builder.intcast(length, builder.cx.type_i64(), false);
502            builder.mul(length_i64, builder.cx.get_const_i64(element_size))
503        }
504        _ => ::rustc_span::macros::bug_impl(None,
    format_args!("unexpected offload size {0:?}", meta.payload_size),
    Location::caller())bug!("unexpected offload size {:?}", meta.payload_size),
505    }
506}
507
508// For each kernel *call*, we now use some of our previous declared globals to move data to and from
509// the gpu. For now, we only handle the data transfer part of it.
510// If two consecutive kernels use the same memory, we still move it to the host and back to the gpu.
511// Since in our frontend users (by default) don't have to specify data transfer, this is something
512// we should optimize in the future! In some cases we can directly zero-allocate on the device and
513// only move data back, or if something is immutable, we might only copy it to the device.
514//
515// Current steps:
516// 0. Alloca some variables for the following steps
517// 1. set insert point before kernel call.
518// 2. generate all the GEPS and stores, to be used in 3)
519// 3. generate __tgt_target_data_begin calls to move data to the GPU
520//
521// unchanged: keep kernel call. Later move the kernel to the GPU
522//
523// 4. set insert point after kernel call.
524// 5. generate all the GEPS and stores, to be used in 6)
525// 6. generate __tgt_target_data_end calls to move data from the GPU
526pub(crate) fn gen_call_handling<'ll, 'tcx>(
527    builder: &mut Builder<'_, 'll, 'tcx>,
528    offload_data: &OffloadKernelGlobals<'ll>,
529    args: &[&'ll Value],
530    types: &[&Type],
531    metadata: &[OffloadMetadata],
532    offload_globals: &OffloadGlobals<'ll>,
533    offload_dims: &OffloadKernelDims<'ll>,
534    dyn_cache: &'ll Value,
535    device_id: &'ll Value,
536) {
537    let cx = builder.cx;
538    let OffloadKernelGlobals {
539        offload_sizes,
540        memtransfer_begin,
541        memtransfer_kernel,
542        memtransfer_end,
543        region_id,
544    } = offload_data;
545    let OffloadKernelDims { num_workgroups, threads_per_block, workgroup_dims, thread_dims } =
546        offload_dims;
547
548    let has_dynamic = metadata.iter().any(|m| !#[allow(non_exhaustive_omitted_patterns)] match m.payload_size {
    OffloadSize::Static(_) => true,
    _ => false,
}matches!(m.payload_size, OffloadSize::Static(_)));
549
550    let tgt_decl = offload_globals.launcher_fn;
551    let tgt_target_kernel_ty = offload_globals.launcher_ty;
552
553    let tgt_kernel_decl = offload_globals.kernel_args_ty;
554    let begin_mapper_decl = offload_globals.begin_mapper;
555    let end_mapper_decl = offload_globals.end_mapper;
556    let fn_ty = offload_globals.mapper_fn_ty;
557
558    let num_args = types.len() as u64;
559    let bb = builder.llbb();
560
561    // Step 0)
562    unsafe {
563        llvm::LLVMRustPositionBuilderPastAllocas(&builder.llbuilder, builder.llfn());
564    }
565
566    let ty = cx.type_array(cx.type_ptr(), num_args);
567    // Baseptr are just the input pointer to the kernel, stored in a local alloca
568    let a1 = builder.direct_alloca(ty, Align::EIGHT, ".offload_baseptrs");
569    // Ptrs are the result of a gep into the baseptr, at least for our trivial types.
570    let a2 = builder.direct_alloca(ty, Align::EIGHT, ".offload_ptrs");
571    // These represent the sizes in bytes, e.g. the entry for `&[f64; 16]` will be 8*16.
572    let ty2 = cx.type_array(cx.type_i64(), num_args);
573
574    let a4 = if has_dynamic {
575        let alloc = builder.direct_alloca(ty2, Align::EIGHT, ".offload_sizes");
576
577        builder.memcpy(
578            alloc,
579            Align::EIGHT,
580            offload_sizes,
581            Align::EIGHT,
582            cx.get_const_i64(8 * args.len() as u64),
583            MemFlags::empty(),
584            None,
585        );
586
587        alloc
588    } else {
589        offload_sizes
590    };
591
592    //%kernel_args = alloca %struct.__tgt_kernel_arguments, align 8
593    let a5 = builder.direct_alloca(tgt_kernel_decl, Align::EIGHT, "kernel_args");
594
595    // Step 1)
596    unsafe {
597        llvm::LLVMPositionBuilderAtEnd(&builder.llbuilder, bb);
598    }
599
600    // Now we allocate once per function param, a copy to be passed to one of our maps.
601    let mut vals = ::alloc::vec::Vec::new()vec![];
602    let mut geps = ::alloc::vec::Vec::new()vec![];
603    let i32_0 = cx.get_const_i32(0);
604    for &v in args {
605        let ty = cx.val_ty(v);
606        let ty_kind = cx.type_kind(ty);
607        let (base_val, gep_base) = match ty_kind {
608            TypeKind::Pointer => (v, v),
609            TypeKind::Half | TypeKind::Float | TypeKind::Double | TypeKind::Integer => {
610                // FIXME(Sa4dUs): check for `f128` support, latest NVIDIA cards support it
611                let num_bits = scalar_width(cx, ty);
612
613                let bb = builder.llbb();
614                unsafe {
615                    llvm::LLVMRustPositionBuilderPastAllocas(builder.llbuilder, builder.llfn());
616                }
617                let addr = builder.direct_alloca(cx.type_i64(), Align::EIGHT, "addr");
618                unsafe {
619                    llvm::LLVMPositionBuilderAtEnd(builder.llbuilder, bb);
620                }
621
622                let cast = builder.bitcast(v, cx.type_ix(num_bits));
623                let value = builder.zext(cast, cx.type_i64());
624                builder.store(value, addr, Align::EIGHT);
625                (value, addr)
626            }
627            other => ::rustc_span::macros::bug_impl(None,
    format_args!("offload does not support {0:?}", other), Location::caller())bug!("offload does not support {other:?}"),
628        };
629
630        let gep = builder.inbounds_gep(cx.type_f32(), gep_base, &[i32_0]);
631
632        vals.push(base_val);
633        geps.push(gep);
634    }
635
636    for i in 0..num_args {
637        let idx = cx.get_const_i32(i);
638        let gep1 = builder.inbounds_gep(ty, a1, &[i32_0, idx]);
639        builder.store(vals[i as usize], gep1, Align::EIGHT);
640        let gep2 = builder.inbounds_gep(ty, a2, &[i32_0, idx]);
641        builder.store(geps[i as usize], gep2, Align::EIGHT);
642
643        if !#[allow(non_exhaustive_omitted_patterns)] match metadata[i as
                usize].payload_size {
    OffloadSize::Static(_) => true,
    _ => false,
}matches!(metadata[i as usize].payload_size, OffloadSize::Static(_)) {
644            let gep3 = builder.inbounds_gep(ty2, a4, &[i32_0, idx]);
645            let size_val = get_runtime_size(builder, args, i as usize, &metadata[i as usize]);
646            builder.store(size_val, gep3, Align::EIGHT);
647        }
648    }
649
650    // For now we have a very simplistic indexing scheme into our
651    // offload_{baseptrs,ptrs,sizes}. We will probably improve this along with our gpu frontend pr.
652    fn get_geps<'ll, 'tcx>(
653        builder: &mut Builder<'_, 'll, 'tcx>,
654        ty: &'ll Type,
655        ty2: &'ll Type,
656        a1: &'ll Value,
657        a2: &'ll Value,
658        a4: &'ll Value,
659        is_dynamic: bool,
660    ) -> [&'ll Value; 3] {
661        let cx = builder.cx;
662        let i32_0 = cx.get_const_i32(0);
663
664        let gep1 = builder.inbounds_gep(ty, a1, &[i32_0, i32_0]);
665        let gep2 = builder.inbounds_gep(ty, a2, &[i32_0, i32_0]);
666        let gep3 = if is_dynamic { builder.inbounds_gep(ty2, a4, &[i32_0, i32_0]) } else { a4 };
667        [gep1, gep2, gep3]
668    }
669
670    fn generate_mapper_call<'ll, 'tcx>(
671        builder: &mut Builder<'_, 'll, 'tcx>,
672        geps: [&'ll Value; 3],
673        o_type: &'ll Value,
674        fn_to_call: &'ll Value,
675        fn_ty: &'ll Type,
676        num_args: u64,
677        s_ident_t: &'ll Value,
678    ) {
679        let cx = builder.cx;
680        let nullptr = cx.const_null(cx.type_ptr());
681        let i64_max = cx.get_const_i64(u64::MAX);
682        let num_args = cx.get_const_i32(num_args);
683        let args =
684            ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
        [s_ident_t, i64_max, num_args, geps[0], geps[1], geps[2], o_type,
                nullptr, nullptr]))vec![s_ident_t, i64_max, num_args, geps[0], geps[1], geps[2], o_type, nullptr, nullptr];
685        builder.call(fn_ty, None, None, fn_to_call, ReturnSlot::Direct, &args, None, None);
686    }
687
688    // Step 2)
689    let s_ident_t = offload_globals.ident_t_global;
690    let geps = get_geps(builder, ty, ty2, a1, a2, a4, has_dynamic);
691    generate_mapper_call(
692        builder,
693        geps,
694        memtransfer_begin,
695        begin_mapper_decl,
696        fn_ty,
697        num_args,
698        s_ident_t,
699    );
700    let values = KernelArgsTy::new(
701        &cx,
702        num_args,
703        memtransfer_kernel,
704        geps,
705        workgroup_dims,
706        thread_dims,
707        dyn_cache,
708    );
709
710    // Step 3)
711    // Here we fill the KernelArgsTy, see the documentation above
712    for (i, value) in values.iter().enumerate() {
713        let ptr = builder.inbounds_gep(tgt_kernel_decl, a5, &[i32_0, cx.get_const_i32(i as u64)]);
714        let name = std::ffi::CString::new(value.1).unwrap();
715        llvm::set_value_name(ptr, &name.as_bytes());
716
717        builder.store(value.2, ptr, value.0);
718    }
719
720    let device_id = builder.sext(device_id, cx.type_i64());
721    let args = ::alloc::boxed::box_assume_init_into_vec_unsafe(::alloc::intrinsics::write_box_via_move(::alloc::boxed::Box::new_uninit(),
        [s_ident_t, device_id, num_workgroups, threads_per_block, region_id,
                a5]))vec![s_ident_t, device_id, num_workgroups, threads_per_block, region_id, a5];
722    builder.call(tgt_target_kernel_ty, None, None, tgt_decl, ReturnSlot::Direct, &args, None, None);
723    // %41 = call i32 @__tgt_target_kernel(ptr @1, i64 -1, i32 2097152, i32 256, ptr @.kernel_1.region_id, ptr %kernel_args)
724
725    // Step 4)
726    let geps = get_geps(builder, ty, ty2, a1, a2, a4, has_dynamic);
727    generate_mapper_call(
728        builder,
729        geps,
730        memtransfer_end,
731        end_mapper_decl,
732        fn_ty,
733        num_args,
734        s_ident_t,
735    );
736}