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
20pub(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 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
112fn 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
127pub(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
141pub(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 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 }
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 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
213struct KernelArgsTy {
215 }
237
238impl KernelArgsTy {
239 const OFFLOAD_VERSION: u64 = 3;
240 const FLAGS: u64 = 1 << 6; 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 (eight, "ArgNames", cx.const_null(cx.type_ptr())), (eight, "ArgMappers", cx.const_null(cx.type_ptr())), (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#[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
356pub(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 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 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 _ => 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 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 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
508pub(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 unsafe {
563 llvm::LLVMRustPositionBuilderPastAllocas(&builder.llbuilder, builder.llfn());
564 }
565
566 let ty = cx.type_array(cx.type_ptr(), num_args);
567 let a1 = builder.direct_alloca(ty, Align::EIGHT, ".offload_baseptrs");
569 let a2 = builder.direct_alloca(ty, Align::EIGHT, ".offload_ptrs");
571 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 let a5 = builder.direct_alloca(tgt_kernel_decl, Align::EIGHT, "kernel_args");
594
595 unsafe {
597 llvm::LLVMPositionBuilderAtEnd(&builder.llbuilder, bb);
598 }
599
600 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 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 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 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 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 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}