333 lines · cpp
1//===- toyc.cpp - The Toy Compiler ----------------------------------------===//2//3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.4// See https://llvm.org/LICENSE.txt for license information.5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception6//7//===----------------------------------------------------------------------===//8//9// This file implements the entry point for the Toy compiler.10//11//===----------------------------------------------------------------------===//12 13#include "mlir/Dialect/Func/Extensions/AllExtensions.h"14#include "mlir/Dialect/Func/IR/FuncOps.h"15#include "mlir/Dialect/LLVMIR/LLVMDialect.h"16#include "mlir/Dialect/LLVMIR/Transforms/InlinerInterfaceImpl.h"17#include "toy/AST.h"18#include "toy/Dialect.h"19#include "toy/Lexer.h"20#include "toy/MLIRGen.h"21#include "toy/Parser.h"22#include "toy/Passes.h"23 24#include "mlir/Dialect/Affine/Passes.h"25#include "mlir/Dialect/LLVMIR/Transforms/Passes.h"26#include "mlir/ExecutionEngine/ExecutionEngine.h"27#include "mlir/ExecutionEngine/OptUtils.h"28#include "mlir/IR/AsmState.h"29#include "mlir/IR/BuiltinOps.h"30#include "mlir/IR/MLIRContext.h"31#include "mlir/IR/Verifier.h"32#include "mlir/InitAllDialects.h"33#include "mlir/Parser/Parser.h"34#include "mlir/Pass/PassManager.h"35#include "mlir/Target/LLVMIR/Dialect/Builtin/BuiltinToLLVMIRTranslation.h"36#include "mlir/Target/LLVMIR/Dialect/LLVMIR/LLVMToLLVMIRTranslation.h"37#include "mlir/Target/LLVMIR/Export.h"38#include "mlir/Transforms/Passes.h"39 40#include "llvm/ADT/StringRef.h"41#include "llvm/ExecutionEngine/Orc/JITTargetMachineBuilder.h"42#include "llvm/IR/Module.h"43#include "llvm/Support/CommandLine.h"44#include "llvm/Support/ErrorOr.h"45#include "llvm/Support/MemoryBuffer.h"46#include "llvm/Support/SourceMgr.h"47#include "llvm/Support/TargetSelect.h"48#include "llvm/Support/raw_ostream.h"49#include <cassert>50#include <memory>51#include <string>52#include <system_error>53#include <utility>54 55using namespace toy;56namespace cl = llvm::cl;57 58static cl::opt<std::string> inputFilename(cl::Positional,59 cl::desc("<input toy file>"),60 cl::init("-"),61 cl::value_desc("filename"));62 63namespace {64enum InputType { Toy, MLIR };65} // namespace66static cl::opt<enum InputType> inputType(67 "x", cl::init(Toy), cl::desc("Decided the kind of output desired"),68 cl::values(clEnumValN(Toy, "toy", "load the input file as a Toy source.")),69 cl::values(clEnumValN(MLIR, "mlir",70 "load the input file as an MLIR file")));71 72namespace {73enum Action {74 None,75 DumpAST,76 DumpMLIR,77 DumpMLIRAffine,78 DumpMLIRLLVM,79 DumpLLVMIR,80 RunJIT81};82} // namespace83static cl::opt<enum Action> emitAction(84 "emit", cl::desc("Select the kind of output desired"),85 cl::values(clEnumValN(DumpAST, "ast", "output the AST dump")),86 cl::values(clEnumValN(DumpMLIR, "mlir", "output the MLIR dump")),87 cl::values(clEnumValN(DumpMLIRAffine, "mlir-affine",88 "output the MLIR dump after affine lowering")),89 cl::values(clEnumValN(DumpMLIRLLVM, "mlir-llvm",90 "output the MLIR dump after llvm lowering")),91 cl::values(clEnumValN(DumpLLVMIR, "llvm", "output the LLVM IR dump")),92 cl::values(93 clEnumValN(RunJIT, "jit",94 "JIT the code and run it by invoking the main function")));95 96static cl::opt<bool> enableOpt("opt", cl::desc("Enable optimizations"));97 98/// Returns a Toy AST resulting from parsing the file or a nullptr on error.99static std::unique_ptr<toy::ModuleAST>100parseInputFile(llvm::StringRef filename) {101 llvm::ErrorOr<std::unique_ptr<llvm::MemoryBuffer>> fileOrErr =102 llvm::MemoryBuffer::getFileOrSTDIN(filename);103 if (std::error_code ec = fileOrErr.getError()) {104 llvm::errs() << "Could not open input file: " << ec.message() << "\n";105 return nullptr;106 }107 auto buffer = fileOrErr.get()->getBuffer();108 LexerBuffer lexer(buffer.begin(), buffer.end(), std::string(filename));109 Parser parser(lexer);110 return parser.parseModule();111}112 113static int loadMLIR(mlir::MLIRContext &context,114 mlir::OwningOpRef<mlir::ModuleOp> &module) {115 // Handle '.toy' input to the compiler.116 if (inputType != InputType::MLIR &&117 !llvm::StringRef(inputFilename).ends_with(".mlir")) {118 auto moduleAST = parseInputFile(inputFilename);119 if (!moduleAST)120 return 6;121 module = mlirGen(context, *moduleAST);122 return !module ? 1 : 0;123 }124 125 // Otherwise, the input is '.mlir'.126 llvm::ErrorOr<std::unique_ptr<llvm::MemoryBuffer>> fileOrErr =127 llvm::MemoryBuffer::getFileOrSTDIN(inputFilename);128 if (std::error_code ec = fileOrErr.getError()) {129 llvm::errs() << "Could not open input file: " << ec.message() << "\n";130 return -1;131 }132 133 // Parse the input mlir.134 llvm::SourceMgr sourceMgr;135 sourceMgr.AddNewSourceBuffer(std::move(*fileOrErr), llvm::SMLoc());136 module = mlir::parseSourceFile<mlir::ModuleOp>(sourceMgr, &context);137 if (!module) {138 llvm::errs() << "Error can't load file " << inputFilename << "\n";139 return 3;140 }141 return 0;142}143 144static int loadAndProcessMLIR(mlir::MLIRContext &context,145 mlir::OwningOpRef<mlir::ModuleOp> &module) {146 if (int error = loadMLIR(context, module))147 return error;148 149 mlir::PassManager pm(module.get()->getName());150 // Apply any generic pass manager command line options and run the pipeline.151 if (mlir::failed(mlir::applyPassManagerCLOptions(pm)))152 return 4;153 154 // Check to see what granularity of MLIR we are compiling to.155 bool isLoweringToAffine = emitAction >= Action::DumpMLIRAffine;156 bool isLoweringToLLVM = emitAction >= Action::DumpMLIRLLVM;157 158 if (enableOpt || isLoweringToAffine) {159 // Inline all functions into main and then delete them.160 pm.addPass(mlir::createInlinerPass());161 162 // Now that there is only one function, we can infer the shapes of each of163 // the operations.164 mlir::OpPassManager &optPM = pm.nest<mlir::toy::FuncOp>();165 optPM.addPass(mlir::toy::createShapeInferencePass());166 optPM.addPass(mlir::createCanonicalizerPass());167 optPM.addPass(mlir::createCSEPass());168 }169 170 if (isLoweringToAffine) {171 // Partially lower the toy dialect.172 pm.addPass(mlir::toy::createLowerToAffinePass());173 174 // Add a few cleanups post lowering.175 mlir::OpPassManager &optPM = pm.nest<mlir::func::FuncOp>();176 optPM.addPass(mlir::createCanonicalizerPass());177 optPM.addPass(mlir::createCSEPass());178 179 // Add optimizations if enabled.180 if (enableOpt) {181 optPM.addPass(mlir::affine::createLoopFusionPass());182 optPM.addPass(mlir::affine::createAffineScalarReplacementPass());183 }184 }185 186 if (isLoweringToLLVM) {187 // Finish lowering the toy IR to the LLVM dialect.188 pm.addPass(mlir::toy::createLowerToLLVMPass());189 // This is necessary to have line tables emitted and basic190 // debugger working. In the future we will add proper debug information191 // emission directly from our frontend.192 pm.addPass(mlir::LLVM::createDIScopeForLLVMFuncOpPass());193 }194 195 if (mlir::failed(pm.run(*module)))196 return 4;197 return 0;198}199 200static int dumpAST() {201 if (inputType == InputType::MLIR) {202 llvm::errs() << "Can't dump a Toy AST when the input is MLIR\n";203 return 5;204 }205 206 auto moduleAST = parseInputFile(inputFilename);207 if (!moduleAST)208 return 1;209 210 dump(*moduleAST);211 return 0;212}213 214static int dumpLLVMIR(mlir::ModuleOp module) {215 // Register the translation to LLVM IR with the MLIR context.216 mlir::registerBuiltinDialectTranslation(*module->getContext());217 mlir::registerLLVMDialectTranslation(*module->getContext());218 219 // Convert the module to LLVM IR in a new LLVM IR context.220 llvm::LLVMContext llvmContext;221 auto llvmModule = mlir::translateModuleToLLVMIR(module, llvmContext);222 if (!llvmModule) {223 llvm::errs() << "Failed to emit LLVM IR\n";224 return -1;225 }226 227 // Initialize LLVM targets.228 llvm::InitializeNativeTarget();229 llvm::InitializeNativeTargetAsmPrinter();230 231 // Configure the LLVM Module232 auto tmBuilderOrError = llvm::orc::JITTargetMachineBuilder::detectHost();233 if (!tmBuilderOrError) {234 llvm::errs() << "Could not create JITTargetMachineBuilder\n";235 return -1;236 }237 238 auto tmOrError = tmBuilderOrError->createTargetMachine();239 if (!tmOrError) {240 llvm::errs() << "Could not create TargetMachine\n";241 return -1;242 }243 mlir::ExecutionEngine::setupTargetTripleAndDataLayout(llvmModule.get(),244 tmOrError.get().get());245 246 /// Optionally run an optimization pipeline over the llvm module.247 auto optPipeline = mlir::makeOptimizingTransformer(248 /*optLevel=*/enableOpt ? 3 : 0, /*sizeLevel=*/0,249 /*targetMachine=*/nullptr);250 if (auto err = optPipeline(llvmModule.get())) {251 llvm::errs() << "Failed to optimize LLVM IR " << err << "\n";252 return -1;253 }254 llvm::errs() << *llvmModule << "\n";255 return 0;256}257 258static int runJit(mlir::ModuleOp module) {259 // Initialize LLVM targets.260 llvm::InitializeNativeTarget();261 llvm::InitializeNativeTargetAsmPrinter();262 263 // Register the translation from MLIR to LLVM IR, which must happen before we264 // can JIT-compile.265 mlir::registerBuiltinDialectTranslation(*module->getContext());266 mlir::registerLLVMDialectTranslation(*module->getContext());267 268 // An optimization pipeline to use within the execution engine.269 auto optPipeline = mlir::makeOptimizingTransformer(270 /*optLevel=*/enableOpt ? 3 : 0, /*sizeLevel=*/0,271 /*targetMachine=*/nullptr);272 273 // Create an MLIR execution engine. The execution engine eagerly JIT-compiles274 // the module.275 mlir::ExecutionEngineOptions engineOptions;276 engineOptions.transformer = optPipeline;277 auto maybeEngine = mlir::ExecutionEngine::create(module, engineOptions);278 assert(maybeEngine && "failed to construct an execution engine");279 auto &engine = maybeEngine.get();280 281 // Invoke the JIT-compiled function.282 auto invocationResult = engine->invokePacked("main");283 if (invocationResult) {284 llvm::errs() << "JIT invocation failed\n";285 return -1;286 }287 288 return 0;289}290 291int main(int argc, char **argv) {292 // Register any command line options.293 mlir::registerAsmPrinterCLOptions();294 mlir::registerMLIRContextCLOptions();295 mlir::registerPassManagerCLOptions();296 297 cl::ParseCommandLineOptions(argc, argv, "toy compiler\n");298 299 if (emitAction == Action::DumpAST)300 return dumpAST();301 302 // If we aren't dumping the AST, then we are compiling with/to MLIR.303 mlir::DialectRegistry registry;304 mlir::func::registerAllExtensions(registry);305 mlir::LLVM::registerInlinerInterface(registry);306 307 mlir::MLIRContext context(registry);308 // Load our Dialect in this MLIR Context.309 context.getOrLoadDialect<mlir::toy::ToyDialect>();310 311 mlir::OwningOpRef<mlir::ModuleOp> module;312 if (int error = loadAndProcessMLIR(context, module))313 return error;314 315 // If we aren't exporting to non-mlir, then we are done.316 bool isOutputingMLIR = emitAction <= Action::DumpMLIRLLVM;317 if (isOutputingMLIR) {318 module->dump();319 return 0;320 }321 322 // Check to see if we are compiling to LLVM IR.323 if (emitAction == Action::DumpLLVMIR)324 return dumpLLVMIR(*module);325 326 // Otherwise, we must be running the jit.327 if (emitAction == Action::RunJIT)328 return runJit(*module);329 330 llvm::errs() << "No action specified (parsing only?), use -emit=<action>\n";331 return -1;332}333