106106#include " comet/Conversion/PrepareGpuHost/PrepareGpuHostPass.h"
107107#include " comet/Conversion/BlockedGpuToTriton/BlockedGpuToTriton.h"
108108#include " comet/Conversion/GpuToBlockedGpu/GpuToBlockedGpu.h"
109- #include " comet/Conversion/ParallelLoopsToGpu/ParallelLoopsToGpu .h"
109+ #include " comet/Conversion/ForallToGpu/ForallToGpu .h"
110110#include " comet/Conversion/TritonToHIP/TritonToHIPPass.h"
111111#include " triton/Dialect/TritonGPU/IR/Dialect.h"
112112#include " mlir/Conversion/SCFToGPU/SCFToGPUPass.h"
@@ -340,7 +340,7 @@ std::unique_ptr<tensorAlgebra::ModuleAST> parseInputFile(llvm::StringRef filenam
340340}
341341
342342int loadMLIR (mlir::MLIRContext &context,
343- mlir::OwningOpRef<mlir::ModuleOp> &module )
343+ mlir::OwningOpRef<mlir::ModuleOp> &module , bool useI64 )
344344{
345345 // / Handle '.ta' input to the compiler.
346346 if (inputType != InputType::MLIR &&
@@ -349,7 +349,7 @@ int loadMLIR(mlir::MLIRContext &context,
349349 auto moduleAST = parseInputFile (inputFilename);
350350 if (!moduleAST)
351351 return 6 ;
352- module = mlirGen (context, *moduleAST);
352+ module = mlirGen (context, *moduleAST, useI64 );
353353 return !module ? 1 : 0 ;
354354 }
355355
@@ -375,7 +375,7 @@ int loadMLIR(mlir::MLIRContext &context,
375375}
376376
377377int loadAndProcessMLIR (mlir::MLIRContext &context,
378- mlir::OwningOpRef<mlir::ModuleOp> &module )
378+ mlir::OwningOpRef<mlir::ModuleOp> &module , bool useI64 )
379379{
380380#ifdef ENABLE_GPU_TARGET
381381 bool emitTriton_ = emitTriton && CodegenTarget == TargetDevice::GPU ;
@@ -389,7 +389,7 @@ int loadAndProcessMLIR(mlir::MLIRContext &context,
389389 tensorAlgebra::debugOptions.insert (" debug-ta-labels-alphabet-order" );
390390 }
391391 // / end Load debug options
392- if (int error = loadMLIR (context, module ))
392+ if (int error = loadMLIR (context, module , useI64 ))
393393 return error;
394394
395395 mlir::PassManager pm (module .get ()->getName ());
@@ -596,9 +596,19 @@ int loadAndProcessMLIR(mlir::MLIRContext &context,
596596 // / Blanket-convert any remaining affine ops if any remain.
597597 pm.addPass (mlir::createLowerAffinePass ());
598598 // / Convert SCF to CF (always needed).
599- pm.addPass (mlir::createForallToParallelLoopPass ());
600- pm.addPass (mlir::createLoopInvariantCodeMotionPass ());
601- pm.addPass (mlir::createCanonicalizerPass ());
599+ #ifdef ENABLE_GPU_TARGET
600+ if (CodegenTarget != TargetDevice::GPU )
601+ {
602+ pm.addPass (mlir::createForallToParallelLoopPass ());
603+ pm.addPass (mlir::createLoopInvariantCodeMotionPass ());
604+ pm.addPass (mlir::createCanonicalizerPass ());
605+ }
606+ #else
607+ pm.addPass (mlir::createForallToParallelLoopPass ());
608+ pm.addPass (mlir::createLoopInvariantCodeMotionPass ());
609+ pm.addPass (mlir::createCanonicalizerPass ());
610+ #endif
611+
602612
603613#ifndef ENABLE_GPU_TARGET
604614 [[maybe_unused]] bool IsLoweringToTriton = false ;
@@ -615,7 +625,7 @@ int loadAndProcessMLIR(mlir::MLIRContext &context,
615625 #ifdef ENABLE_GPU_TARGET
616626 if (CodegenTarget == TargetDevice::GPU )
617627 {
618- pm.addNestedPass <mlir::func::FuncOp>(mlir::comet::createConvertParallelLoopsToGpuPass (GPUBlockSizeX, GPUBlockSizeY, GPUBlockSizeR));
628+ pm.addNestedPass <mlir::func::FuncOp>(mlir::comet::createConvertForallToGpuPass (GPUBlockSizeX, GPUBlockSizeY, GPUBlockSizeR));
619629 }
620630 #endif
621631
@@ -804,8 +814,16 @@ int main(int argc, char **argv)
804814 context.loadDialect <mlir::index::IndexDialect>();
805815
806816 mlir::OwningOpRef<mlir::ModuleOp> module ;
817+ bool useI64 = true ;
818+ #ifdef ENABLE_GPU_TARGET
819+ if (CodegenTarget == TargetDevice::GPU )
820+ {
821+ useI64 = false ;
822+ }
823+ #endif
824+
807825
808- if (int error = loadAndProcessMLIR (context, module ))
826+ if (int error = loadAndProcessMLIR (context, module , useI64 ))
809827 return error;
810828
811829 // / If we aren't exporting to non-mlir, then we are done.
0 commit comments