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 a910c1ff69..691acaef8a 100644 --- a/quadrants/codegen/amdgpu/codegen_amdgpu.cpp +++ b/quadrants/codegen/amdgpu/codegen_amdgpu.cpp @@ -543,6 +543,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; @@ -584,6 +593,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