diff --git a/Dockerfile.rocm b/Dockerfile.rocm
index c1e48190c1..76f6b22f9d 100644
--- a/Dockerfile.rocm
+++ b/Dockerfile.rocm
@@ -29,14 +29,14 @@ RUN uv venv -p ${PY_VERSION} /opt/venv
ENV VIRTUAL_ENV=/opt/venv
ENV PATH="/opt/venv/bin:$PATH"
-# rocm and torch MUST be installed together or torch nightlies will
+# rocm and torch MUST be installed together or torch nightlies will
# mess up the rocm version
RUN uv pip install pip cmake pybind11 build ninja scikit-build-core setuptools-scm numpy pytest && \
uv pip install --index-url https://repo.amd.com/rocm/whl/${GFX_FAMILY}/ "rocm[libraries,devel]"==${THE_ROCK_VERSION} torch==${TORCH_VERSION} && \
rocm-sdk init
# Setup SRC
-ENV QUADRANTS_SRC_DIR=/src/quadrants/
+ENV QUADRANTS_SRC_DIR=/src/quadrants/
COPY . ${QUADRANTS_SRC_DIR}
RUN git config --global --add safe.directory ${QUADRANTS_SRC_DIR}
diff --git a/python/quadrants/lang/_func_base.py b/python/quadrants/lang/_func_base.py
index 15067182f3..c285ea072b 100644
--- a/python/quadrants/lang/_func_base.py
+++ b/python/quadrants/lang/_func_base.py
@@ -1,6 +1,5 @@
import ast
import inspect
-import math
import sys
import textwrap
import types
diff --git a/python/quadrants/lang/ast/ast_transformer.py b/python/quadrants/lang/ast/ast_transformer.py
index 79c639761e..64acc8d79e 100644
--- a/python/quadrants/lang/ast/ast_transformer.py
+++ b/python/quadrants/lang/ast/ast_transformer.py
@@ -905,13 +905,15 @@ def build_static_for(ctx: ASTTransformerFuncContext, node: ast.For, is_grouped:
@staticmethod
def build_range_for(ctx: ASTTransformerFuncContext, node: ast.For) -> None:
+ if len(node.iter.args) not in [1, 2, 3]:
+ raise QuadrantsSyntaxError(f"Range should have 1, 2, or 3 arguments, found {len(node.iter.args)}")
+ if len(node.iter.args) == 3:
+ return ASTTransformer.build_strided_range_for(ctx, node)
with ctx.variable_scope_guard():
loop_name = node.target.id
ctx.check_loop_var(loop_name)
loop_var = expr.Expr(ctx.ast_builder.make_id_expr(""))
ctx.create_variable(loop_name, loop_var)
- if len(node.iter.args) not in [1, 2]:
- raise QuadrantsSyntaxError(f"Range should have 1 or 2 arguments, found {len(node.iter.args)}")
if len(node.iter.args) == 2:
begin_expr = expr.Expr(build_stmt(ctx, node.iter.args[0]))
end_expr = expr.Expr(build_stmt(ctx, node.iter.args[1]))
@@ -940,6 +942,50 @@ def build_range_for(ctx: ASTTransformerFuncContext, node: ast.For) -> None:
ctx.ast_builder.end_frontend_range_for()
return None
+ @staticmethod
+ def build_strided_range_for(ctx, node):
+ """Desugar `for i in range(start, stop, step)` into a while loop.
+
+ The Quadrants IR does not natively support a step parameter in
+ range-for loops. We lower `range(start, stop, step)` into::
+
+ i = start
+ while i < stop: # (or i > stop when step < 0)
+
+ i = i + step
+ """
+ with ctx.variable_scope_guard():
+ loop_name = node.target.id
+
+ begin_expr = expr.Expr(build_stmt(ctx, node.iter.args[0]))
+ end_expr = expr.Expr(build_stmt(ctx, node.iter.args[1]))
+ step_expr = expr.Expr(build_stmt(ctx, node.iter.args[2]))
+
+ begin = qd_ops.cast(begin_expr, primitive_types.i32)
+ end = qd_ops.cast(end_expr, primitive_types.i32)
+ step = qd_ops.cast(step_expr, primitive_types.i32)
+
+ loop_var = impl.expr_init(begin)
+ ctx.create_variable(loop_name, loop_var)
+
+ with ctx.loop_scope_guard():
+ stmt_dbg_info = _qd_core.DebugInfo(ctx.get_pos_info(node))
+ ctx.ast_builder.begin_frontend_while(expr.Expr(1, dtype=primitive_types.i32).ptr, stmt_dbg_info)
+
+ cond = loop_var < end
+ impl.begin_frontend_if(ctx.ast_builder, cond, stmt_dbg_info)
+ ctx.ast_builder.begin_frontend_if_true()
+ ctx.ast_builder.pop_scope()
+ ctx.ast_builder.begin_frontend_if_false()
+ ctx.ast_builder.insert_break_stmt(stmt_dbg_info)
+ ctx.ast_builder.pop_scope()
+
+ build_stmts(ctx, node.body)
+
+ loop_var._assign(loop_var + step)
+ ctx.ast_builder.pop_scope()
+ return None
+
@staticmethod
def build_ndrange_for(ctx: ASTTransformerFuncContext, node: ast.For) -> None:
with ctx.variable_scope_guard():
diff --git a/python/quadrants/lang/simt/subgroup.py b/python/quadrants/lang/simt/subgroup.py
index edec8978d8..046dfae7f1 100644
--- a/python/quadrants/lang/simt/subgroup.py
+++ b/python/quadrants/lang/simt/subgroup.py
@@ -147,6 +147,10 @@ def shuffle_xor(value, mask):
pass
+def dpp_swap_pairs(value):
+ return impl.call_internal("subgroupDppSwapPairs", value, with_runtime_context=False)
+
+
def shuffle_up(value, offset):
return impl.call_internal("subgroupShuffleUp", value, offset, with_runtime_context=False)
@@ -188,4 +192,5 @@ def shuffle_down(value, offset):
"shuffle_xor",
"shuffle_up",
"shuffle_down",
+ "dpp_swap_pairs",
]
diff --git a/quadrants/codegen/amdgpu/codegen_amdgpu.cpp b/quadrants/codegen/amdgpu/codegen_amdgpu.cpp
index f0e166e4fc..86ddb7d798 100644
--- a/quadrants/codegen/amdgpu/codegen_amdgpu.cpp
+++ b/quadrants/codegen/amdgpu/codegen_amdgpu.cpp
@@ -72,8 +72,7 @@ class TaskCodeGenAMDGPU : public TaskCodeGenLLVM {
auto base = new llvm::GlobalVariable(
*module, type, false, llvm::GlobalValue::ExternalLinkage, nullptr,
fmt::format("shared_array_t{}_s{}", task_codegen_id, stmt->id),
- nullptr, llvm::GlobalVariable::NotThreadLocal,
- 3 /*addrspace=LDS*/);
+ nullptr, llvm::GlobalVariable::NotThreadLocal, 3 /*addrspace=LDS*/);
base->setAlignment(llvm::MaybeAlign(8));
auto ptr_type = llvm::PointerType::get(type, 0);
llvm_val[stmt] = builder->CreatePointerCast(base, ptr_type);
@@ -314,8 +313,7 @@ class TaskCodeGenAMDGPU : public TaskCodeGenLLVM {
if (input && input->getType()->isPointerTy() &&
input->getType()->getPointerAddressSpace() == 1) {
auto *ptr_as0 = llvm::PointerType::getUnqual(*llvm_context);
- llvm_val[stmt->input_ptr] =
- builder->CreateAddrSpaceCast(input, ptr_as0);
+ llvm_val[stmt->input_ptr] = builder->CreateAddrSpaceCast(input, ptr_as0);
}
TaskCodeGenLLVM::visit(stmt);
llvm_val[stmt->input_ptr] = input;
@@ -368,16 +366,15 @@ class TaskCodeGenAMDGPU : public TaskCodeGenLLVM {
tlctx->get_data_type(stmt->origin->ret_type.ptr_removed());
auto *casted_ptr = builder->CreateBitCast(
origin_ptr, llvm::PointerType::get(origin_pointee_ty, origin_as));
- llvm_val[stmt] = builder->CreateGEP(
- origin_pointee_ty, casted_ptr,
- {tlctx->get_constant(0), llvm_val[stmt->offset]});
+ llvm_val[stmt] =
+ builder->CreateGEP(origin_pointee_ty, casted_ptr,
+ {tlctx->get_constant(0), llvm_val[stmt->offset]});
} else {
auto *origin_address = builder->CreatePtrToInt(
origin_ptr, llvm::Type::getInt64Ty(*llvm_context));
auto *address_offset = builder->CreateSExt(
llvm_val[stmt->offset], llvm::Type::getInt64Ty(*llvm_context));
- auto *target_address =
- builder->CreateAdd(origin_address, address_offset);
+ auto *target_address = builder->CreateAdd(origin_address, address_offset);
auto pointee_ty = tlctx->get_data_type(stmt->ret_type.ptr_removed());
llvm_val[stmt] = builder->CreateIntToPtr(
target_address, llvm::PointerType::get(pointee_ty, origin_as));
@@ -443,8 +440,8 @@ class TaskCodeGenAMDGPU : public TaskCodeGenLLVM {
// BLS / shared memory buffer allocation
void create_bls_buffer(OffloadedStmt *stmt) {
- auto type = llvm::ArrayType::get(
- llvm::Type::getInt8Ty(*llvm_context), stmt->bls_size);
+ auto type = llvm::ArrayType::get(llvm::Type::getInt8Ty(*llvm_context),
+ stmt->bls_size);
bls_buffer = new llvm::GlobalVariable(
*module, type, false, llvm::GlobalValue::ExternalLinkage, nullptr,
"bls_buffer", nullptr, llvm::GlobalVariable::NotThreadLocal,
@@ -524,6 +521,15 @@ class TaskCodeGenAMDGPU : public TaskCodeGenLLVM {
}
}
+ void visit(InternalFuncStmt *stmt) override {
+ if (stmt->func_name == "subgroupDppSwapPairs") {
+ llvm_val[stmt] = emit_amdgpu_dpp_swap_pairs(llvm_val[stmt->args[0]],
+ stmt->args[0]->ret_type);
+ } else {
+ TaskCodeGenLLVM::visit(stmt);
+ }
+ }
+
void visit(BinaryOpStmt *stmt) override {
auto op = stmt->op_type;
auto ret_quadrants_type = stmt->ret_type;
@@ -565,6 +571,45 @@ class TaskCodeGenAMDGPU : public TaskCodeGenLLVM {
}
private:
+ llvm::Value *emit_amdgpu_dpp_swap_pairs(llvm::Value *value, DataType dt) {
+ auto *i32_ty = llvm::Type::getInt32Ty(*llvm_context);
+ auto *i1_ty = llvm::Type::getInt1Ty(*llvm_context);
+ auto *ctrl = llvm::ConstantInt::get(i32_ty, 0xB1);
+ auto *rmask = llvm::ConstantInt::get(i32_ty, 0xF);
+ auto *bmask = llvm::ConstantInt::get(i32_ty, 0xF);
+ auto *bctrl = llvm::ConstantInt::getFalse(i1_ty);
+
+ auto emit_dpp_32 = [&](llvm::Value *v) -> llvm::Value * {
+ auto *ty = v->getType();
+ return builder->CreateIntrinsic(
+ Intrinsic::amdgcn_update_dpp, {ty},
+ {llvm::Constant::getNullValue(ty), v, ctrl, rmask, bmask, bctrl});
+ };
+
+ if (dt->is_primitive(PrimitiveTypeID::i32) ||
+ dt->is_primitive(PrimitiveTypeID::u32) ||
+ dt->is_primitive(PrimitiveTypeID::f32)) {
+ return emit_dpp_32(value);
+ }
+ if (dt->is_primitive(PrimitiveTypeID::f64) ||
+ dt->is_primitive(PrimitiveTypeID::i64) ||
+ dt->is_primitive(PrimitiveTypeID::u64)) {
+ auto *i64_ty = llvm::Type::getInt64Ty(*llvm_context);
+ auto *i64_val = builder->CreateBitCast(value, i64_ty);
+ auto *lo = builder->CreateTrunc(i64_val, i32_ty);
+ auto *hi = builder->CreateTrunc(builder->CreateLShr(i64_val, 32), i32_ty);
+ lo = emit_dpp_32(lo);
+ hi = emit_dpp_32(hi);
+ auto *result = builder->CreateOr(
+ builder->CreateZExt(lo, i64_ty),
+ builder->CreateShl(builder->CreateZExt(hi, i64_ty), 32));
+ return builder->CreateBitCast(result, value->getType());
+ }
+ QD_ERROR("subgroupDppSwapPairs: unsupported type {} on AMDGPU",
+ data_type_name(dt));
+ return nullptr;
+ }
+
std::tuple get_spmd_info() override {
auto thread_idx = builder->CreateIntrinsic(Intrinsic::amdgcn_workitem_id_x,
ArrayRef{});
diff --git a/quadrants/codegen/llvm/codegen_llvm.cpp b/quadrants/codegen/llvm/codegen_llvm.cpp
index 22311b0c9a..4c55648586 100644
--- a/quadrants/codegen/llvm/codegen_llvm.cpp
+++ b/quadrants/codegen/llvm/codegen_llvm.cpp
@@ -1954,7 +1954,8 @@ void TaskCodeGenLLVM::visit(ExternalPtrStmt *stmt) {
// Indexing array dimensions
linear_index = builder->CreateMul(linear_index, sizes[size_var_index++]);
}
- auto index = builder->CreateSExtOrBitCast(llvm_val[stmt->indices[i]], i64_ty);
+ auto index =
+ builder->CreateSExtOrBitCast(llvm_val[stmt->indices[i]], i64_ty);
linear_index = builder->CreateAdd(linear_index, index);
}
QD_ASSERT(size_var_index == num_indices - num_element_indices);
@@ -2030,9 +2031,8 @@ std::string TaskCodeGenLLVM::init_offloaded_task_function(OffloadedStmt *stmt,
} else {
context_param_type = llvm::PointerType::get(context_ty, 0);
}
- task_function_type =
- llvm::FunctionType::get(llvm::Type::getVoidTy(*llvm_context),
- {context_param_type}, false);
+ task_function_type = llvm::FunctionType::get(
+ llvm::Type::getVoidTy(*llvm_context), {context_param_type}, false);
auto task_kernel_name = fmt::format(
"{}_{}_{}{}", kernel_name, task_codegen_id, stmt->task_name(), suffix);
@@ -2446,6 +2446,16 @@ void TaskCodeGenLLVM::visit(ClearListStmt *stmt) {
}
void TaskCodeGenLLVM::visit(InternalFuncStmt *stmt) {
+ if (stmt->func_name == "subgroupDppSwapPairs") {
+ QD_ERROR(
+ "Internal op \"{}\" requires a GPU backend (AMDGPU or CUDA). "
+ "Wrap the call site with a backend guard such as "
+ "qd.static(backend == gs.amdgpu) so the CPU path never reaches "
+ "it.",
+ stmt->func_name);
+ return;
+ }
+
std::vector args;
if (stmt->with_runtime_context)
diff --git a/quadrants/inc/internal_ops.inc.h b/quadrants/inc/internal_ops.inc.h
index bd8c9c82c1..110bfc5165 100644
--- a/quadrants/inc/internal_ops.inc.h
+++ b/quadrants/inc/internal_ops.inc.h
@@ -27,6 +27,7 @@ PER_INTERNAL_OP(subgroupBarrier)
PER_INTERNAL_OP(subgroupMemoryBarrier)
PER_INTERNAL_OP(subgroupElect)
PER_INTERNAL_OP(subgroupBroadcast)
+PER_INTERNAL_OP(subgroupDppSwapPairs)
PER_INTERNAL_OP(subgroupSize)
PER_INTERNAL_OP(subgroupInvocationId)
PER_INTERNAL_OP(subgroupAdd)
diff --git a/quadrants/ir/type_system.cpp b/quadrants/ir/type_system.cpp
index 85706cff8a..5bce5443d6 100644
--- a/quadrants/ir/type_system.cpp
+++ b/quadrants/ir/type_system.cpp
@@ -378,6 +378,7 @@ void Operations::init_internals() {
PLAIN_OP(subgroupMemoryBarrier, i32_void, false);
PLAIN_OP(subgroupElect, i32, false);
POLY_OP(subgroupBroadcast, false, Signature({}, {ValueT, !u32}, ValueT));
+ POLY_OP(subgroupDppSwapPairs, false, Signature({}, {ValueT}, ValueT));
PLAIN_OP(subgroupSize, i32, false);
PLAIN_OP(subgroupInvocationId, i32, false);
POLY_OP(subgroupAdd, false, Signature({}, {ValueT}, ValueT));
diff --git a/quadrants/rhi/amdgpu/amdgpu_context.cpp b/quadrants/rhi/amdgpu/amdgpu_context.cpp
index 4880cc1fb0..a3f1630099 100644
--- a/quadrants/rhi/amdgpu/amdgpu_context.cpp
+++ b/quadrants/rhi/amdgpu/amdgpu_context.cpp
@@ -190,8 +190,7 @@ void AMDGPUContext::launch(void *func,
bool valid =
offline_cache::try_demangle_name(task_name, primal_task_name, key);
profiler_amdgpu->trace(task_handle, valid ? primal_task_name : task_name,
- func, grid_dim, block_dim,
- dynamic_shared_mem_bytes);
+ func, grid_dim, block_dim, dynamic_shared_mem_bytes);
}
auto context_guard = AMDGPUContext::get_instance().get_guard();
diff --git a/quadrants/rhi/amdgpu/amdgpu_driver_functions.inc.h b/quadrants/rhi/amdgpu/amdgpu_driver_functions.inc.h
index d4cabe73a2..158851621e 100644
--- a/quadrants/rhi/amdgpu/amdgpu_driver_functions.inc.h
+++ b/quadrants/rhi/amdgpu/amdgpu_driver_functions.inc.h
@@ -69,7 +69,11 @@ PER_AMDGPU_FUNCTION(memcpy_device_to_host_async,
std::size_t,
void *);
PER_AMDGPU_FUNCTION(malloc, hipMalloc, void **, std::size_t);
-PER_AMDGPU_FUNCTION(malloc_async_impl, hipMallocAsync, void **, std::size_t, void *);
+PER_AMDGPU_FUNCTION(malloc_async_impl,
+ hipMallocAsync,
+ void **,
+ std::size_t,
+ void *);
PER_AMDGPU_FUNCTION(malloc_managed,
hipMallocManaged,
void **,
diff --git a/quadrants/runtime/amdgpu/jit_amdgpu.cpp b/quadrants/runtime/amdgpu/jit_amdgpu.cpp
index 20f3809a70..e14b21dbf6 100644
--- a/quadrants/runtime/amdgpu/jit_amdgpu.cpp
+++ b/quadrants/runtime/amdgpu/jit_amdgpu.cpp
@@ -17,8 +17,8 @@
namespace quadrants {
namespace lang {
#if defined(QD_WITH_AMDGPU)
-JITModule *JITSessionAMDGPU ::add_module(std::unique_ptr M,
- int max_reg) {
+JITModule *JITSessionAMDGPU::add_module(std::unique_ptr M,
+ int max_reg) {
// HSACo caching
auto cache_key = compute_module_cache_key(M.get());
auto cache_it = hsaco_cache_.find(cache_key);
@@ -50,8 +50,7 @@ std::string JITSessionAMDGPU::compile_module_to_hsaco(
std::unique_ptr &llvm_module) {
static std::once_flag amdgpu_cl_flags;
std::call_once(amdgpu_cl_flags, [] {
- const char *args[] = {"quadrants",
- "-force-vector-interleave=8"};
+ const char *args[] = {"quadrants", "-force-vector-interleave=8"};
llvm::cl::ParseCommandLineOptions(2, args);
});
@@ -118,10 +117,11 @@ std::string JITSessionAMDGPU::compile_module_to_hsaco(
if (CB->getCalledOperand() != &F)
continue;
auto *Caller = CB->getFunction();
- if (Caller && Caller->getCallingConv() == llvm::CallingConv::AMDGPU_KERNEL &&
+ if (Caller &&
+ Caller->getCallingConv() == llvm::CallingConv::AMDGPU_KERNEL &&
Caller->hasFnAttribute("amdgpu-flat-work-group-size")) {
- inherited =
- Caller->getFnAttribute("amdgpu-flat-work-group-size").getValueAsString();
+ inherited = Caller->getFnAttribute("amdgpu-flat-work-group-size")
+ .getValueAsString();
break;
}
}
diff --git a/quadrants/runtime/amdgpu/kernel_launcher.cpp b/quadrants/runtime/amdgpu/kernel_launcher.cpp
index 02b91f1310..000ffd93e3 100644
--- a/quadrants/runtime/amdgpu/kernel_launcher.cpp
+++ b/quadrants/runtime/amdgpu/kernel_launcher.cpp
@@ -74,8 +74,8 @@ void KernelLauncher::launch_offloaded_tasks(
QD_TRACE("Launching kernel {}<<<{}, {}>>>", task.name, task.grid_dim,
task.block_dim);
amdgpu_module->launch(task.name, task.grid_dim, task.block_dim,
- task.dynamic_shared_array_bytes,
- {&ctx.get_context()}, {kRuntimeContextArgSize});
+ task.dynamic_shared_array_bytes, {&ctx.get_context()},
+ {kRuntimeContextArgSize});
}
}
@@ -152,14 +152,14 @@ void KernelLauncher::launch_llvm_kernel(Handle handle,
LaunchContextBuilder::DevAllocType::kNone) {
if (on_amdgpu_device(data_ptr)) {
if (branch_counts) {
- branch_counts->kNone_on_device.fetch_add(
- 1, std::memory_order_relaxed);
+ branch_counts->kNone_on_device.fetch_add(1,
+ std::memory_order_relaxed);
}
device_ptrs[data_ptr_idx] = data_ptr;
} else {
if (branch_counts) {
- branch_counts->kNone_host_copy.fetch_add(
- 1, std::memory_order_relaxed);
+ branch_counts->kNone_host_copy.fetch_add(1,
+ std::memory_order_relaxed);
}
DeviceAllocation devalloc = executor->allocate_memory_on_device(
arr_sz, (uint64 *)device_result_buffer);
diff --git a/quadrants/runtime/llvm/llvm_context.cpp b/quadrants/runtime/llvm/llvm_context.cpp
index 93044da7fc..bec5b19fb3 100644
--- a/quadrants/runtime/llvm/llvm_context.cpp
+++ b/quadrants/runtime/llvm/llvm_context.cpp
@@ -917,8 +917,8 @@ void QuadrantsLLVMContext::mark_function_as_cuda_kernel(llvm::Function *func,
}
}
-void QuadrantsLLVMContext::mark_function_as_amdgpu_kernel(
- llvm::Function *func, int block_dim) {
+void QuadrantsLLVMContext::mark_function_as_amdgpu_kernel(llvm::Function *func,
+ int block_dim) {
func->setCallingConv(llvm::CallingConv::AMDGPU_KERNEL);
if (block_dim > 0) {
// Note: hardcoded wavefront size of 64 matches CDNA3. RDNA in wave32
@@ -927,8 +927,8 @@ void QuadrantsLLVMContext::mark_function_as_amdgpu_kernel(
constexpr int kAmdgpuWavefrontSize = 64;
int min_block_dim = std::max(block_dim, kAmdgpuWavefrontSize);
int max_block_dim = std::max(block_dim, kAmdgpuWavefrontSize);
- std::string size_str = std::to_string(min_block_dim) + "," +
- std::to_string(max_block_dim);
+ std::string size_str =
+ std::to_string(min_block_dim) + "," + std::to_string(max_block_dim);
func->addFnAttr("amdgpu-flat-work-group-size", size_str);
}
}
diff --git a/quadrants/runtime/llvm/llvm_context_pass.h b/quadrants/runtime/llvm/llvm_context_pass.h
index d686c8cc41..c707f48e78 100644
--- a/quadrants/runtime/llvm/llvm_context_pass.h
+++ b/quadrants/runtime/llvm/llvm_context_pass.h
@@ -122,10 +122,12 @@ struct AMDGPUConvertAllocaInstAddressSpacePass : public FunctionPass {
// (scratch), the load/store is left as flat.
struct AMDGPUFlatToGlobalLoadStorePass : public FunctionPass {
static inline char ID{0};
- AMDGPUFlatToGlobalLoadStorePass() : FunctionPass(ID) {}
+ AMDGPUFlatToGlobalLoadStorePass() : FunctionPass(ID) {
+ }
- static bool originatesFromScratch(llvm::Value *ptr,
- llvm::SmallPtrSetImpl &Visited) {
+ static bool originatesFromScratch(
+ llvm::Value *ptr,
+ llvm::SmallPtrSetImpl &Visited) {
auto *origin = ptr->stripPointerCasts();
if (!Visited.insert(origin).second)
return false; // already on the walk path — break the cycle
@@ -205,13 +207,13 @@ struct AMDGPUFlatToGlobalLoadStorePass : public FunctionPass {
for (auto *I : to_convert) {
llvm::IRBuilder<> B(I);
if (auto *LI = llvm::dyn_cast(I)) {
- auto *cast = B.CreateAddrSpaceCast(LI->getPointerOperand(),
- ptr_global_ty);
+ auto *cast =
+ B.CreateAddrSpaceCast(LI->getPointerOperand(), ptr_global_ty);
LI->setOperand(LI->getPointerOperandIndex(), cast);
modified = true;
} else if (auto *SI = llvm::dyn_cast(I)) {
- auto *cast = B.CreateAddrSpaceCast(SI->getPointerOperand(),
- ptr_global_ty);
+ auto *cast =
+ B.CreateAddrSpaceCast(SI->getPointerOperand(), ptr_global_ty);
SI->setOperand(SI->getPointerOperandIndex(), cast);
modified = true;
}
diff --git a/tests/python/test_ast_refactor.py b/tests/python/test_ast_refactor.py
index 21ec12c573..a51ff67cf7 100644
--- a/tests/python/test_ast_refactor.py
+++ b/tests/python/test_ast_refactor.py
@@ -399,15 +399,18 @@ def foo(x: qd.i32):
def test_range_for_three_arguments():
a = qd.field(qd.i32, shape=(10,))
- with pytest.raises(qd.QuadrantsCompilationError, match="Range should have 1 or 2 arguments, found 3"):
-
- @qd.kernel
- def foo(x: qd.i32):
- for i in range(3, 7, 2):
- a[i] = x
+ @qd.kernel
+ def foo(x: qd.i32):
+ for i in range(3, 7, 2):
+ a[i] = x
- x = 5
- foo(x)
+ a.fill(0)
+ foo(5)
+ for i in range(10):
+ if i in (3, 5):
+ assert a[i] == 5
+ else:
+ assert a[i] == 0
@test_utils.test(print_preprocessed_ir=True)
diff --git a/tests/python/test_exception.py b/tests/python/test_exception.py
index df145e155e..32c7e73fb0 100644
--- a/tests/python/test_exception.py
+++ b/tests/python/test_exception.py
@@ -108,9 +108,10 @@ def foo():
@test_utils.test(print_full_traceback=False)
def test_exception_in_node_with_body():
frameinfo = getframeinfo(currentframe())
+
@qd.kernel
def foo():
- for i in range(1, 2, 3):
+ for i in range():
a = 1
b = 1
c = 1
@@ -121,9 +122,8 @@ def foo():
lineno = frameinfo.lineno
file = frameinfo.filename
msg = f"""
-File "{file}", line {lineno + 3}, in foo:
- for i in range(1, 2, 3):
- ^^^^^^^^^^^^^^^^^^^^^^^^
-Range should have 1 or 2 arguments, found 3"""
+File "{file}", line {lineno + 4}, in foo:
+ for i in range():
+ ^^^^^^^^^^^^^^^^^
+Range should have 1, 2, or 3 arguments, found 0"""
assert msg in e.value.args[0]
-