brintos

brintos / llvm-project-archived public Read only

0
0
Text · 11.1 KiB · 031eae2 Raw
346 lines · cpp
1//===- Remarks.cpp - MLIR Remarks -----------------------------------------===//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#include "mlir/IR/Remarks.h"10 11#include "mlir/IR/BuiltinAttributes.h"12#include "mlir/IR/Diagnostics.h"13#include "mlir/IR/Value.h"14 15#include "llvm/ADT/StringExtras.h"16#include "llvm/ADT/StringRef.h"17 18using namespace mlir::remark::detail;19using namespace mlir::remark;20//------------------------------------------------------------------------------21// Remark22//------------------------------------------------------------------------------23 24Remark::Arg::Arg(llvm::StringRef k, Value v) : key(k) {25  llvm::raw_string_ostream os(val);26  os << v;27}28 29Remark::Arg::Arg(llvm::StringRef k, Type t) : key(k) {30  llvm::raw_string_ostream os(val);31  os << t;32}33 34void Remark::insert(llvm::StringRef s) { args.emplace_back(s); }35void Remark::insert(Arg a) { args.push_back(std::move(a)); }36 37// Simple helper to print key=val list (sorted).38static void printArgs(llvm::raw_ostream &os, llvm::ArrayRef<Remark::Arg> args) {39  if (args.empty())40    return;41 42  llvm::SmallVector<Remark::Arg, 8> sorted(args.begin(), args.end());43  llvm::sort(sorted, [](const Remark::Arg &a, const Remark::Arg &b) {44    return a.key < b.key;45  });46 47  for (size_t i = 0; i < sorted.size(); ++i) {48    const auto &a = sorted[i];49    os << a.key << "=";50 51    llvm::StringRef val(a.val);52    bool needsQuote = val.contains(' ') || val.contains(',') ||53                      val.contains('{') || val.contains('}');54    if (needsQuote)55      os << '"' << val << '"';56    else57      os << val;58 59    if (i + 1 < sorted.size())60      os << ", ";61  }62}63 64/// Print the remark to the given output stream.65/// Example output:66// clang-format off67/// [Missed] Category: Loop | Pass:Unroller |  Function=main | Reason="tripCount=4 < threshold=256"68/// [Failure] LoopOptimizer | Reason="failed due to unsupported pattern"69// clang-format on70void Remark::print(llvm::raw_ostream &os, bool printLocation) const {71  // Header: [Type] pass:remarkName72  StringRef type = getRemarkTypeString();73  StringRef categoryName = getCombinedCategoryName();74  StringRef name = remarkName;75 76  os << '[' << type << "] ";77  os << name << " | ";78  if (!categoryName.empty())79    os << "Category:" << categoryName << " | ";80  if (!functionName.empty())81    os << "Function=" << getFunction() << " | ";82 83  if (printLocation) {84    if (auto flc = mlir::dyn_cast<mlir::FileLineColLoc>(getLocation())) {85      os << " @" << flc.getFilename() << ":" << flc.getLine() << ":"86         << flc.getColumn();87    }88  }89 90  printArgs(os, getArgs());91}92 93std::string Remark::getMsg() const {94  std::string s;95  llvm::raw_string_ostream os(s);96  print(os);97  os.flush();98  return s;99}100 101llvm::StringRef Remark::getRemarkTypeString() const {102  switch (remarkKind) {103  case RemarkKind::RemarkUnknown:104    return "Unknown";105  case RemarkKind::RemarkPassed:106    return "Passed";107  case RemarkKind::RemarkMissed:108    return "Missed";109  case RemarkKind::RemarkFailure:110    return "Failure";111  case RemarkKind::RemarkAnalysis:112    return "Analysis";113  }114  llvm_unreachable("Unknown remark kind");115}116 117llvm::remarks::Type Remark::getRemarkType() const {118  switch (remarkKind) {119  case RemarkKind::RemarkUnknown:120    return llvm::remarks::Type::Unknown;121  case RemarkKind::RemarkPassed:122    return llvm::remarks::Type::Passed;123  case RemarkKind::RemarkMissed:124    return llvm::remarks::Type::Missed;125  case RemarkKind::RemarkFailure:126    return llvm::remarks::Type::Failure;127  case RemarkKind::RemarkAnalysis:128    return llvm::remarks::Type::Analysis;129  }130  llvm_unreachable("Unknown remark kind");131}132 133llvm::remarks::Remark Remark::generateRemark() const {134  auto locLambda = [&]() -> llvm::remarks::RemarkLocation {135    if (auto flc = dyn_cast<FileLineColLoc>(getLocation()))136      return {flc.getFilename(), flc.getLine(), flc.getColumn()};137    return {"<unknown file>", 0, 0};138  };139 140  llvm::remarks::Remark r; // The result.141  r.RemarkType = getRemarkType();142  r.RemarkName = getRemarkName();143  // MLIR does not use passes; instead, it has categories and sub-categories.144  r.PassName = getCombinedCategoryName();145  r.FunctionName = getFunction();146  r.Loc = locLambda();147  for (const Remark::Arg &arg : getArgs()) {148    r.Args.emplace_back();149    r.Args.back().Key = arg.key;150    r.Args.back().Val = arg.val;151  }152  return r;153}154 155//===----------------------------------------------------------------------===//156// InFlightRemark157//===----------------------------------------------------------------------===//158 159InFlightRemark::~InFlightRemark() {160  if (remark && owner)161    owner->report(std::move(*remark));162  owner = nullptr;163}164 165//===----------------------------------------------------------------------===//166// Remark Engine167//===----------------------------------------------------------------------===//168 169template <typename RemarkT, typename... Args>170InFlightRemark RemarkEngine::makeRemark(Args &&...args) {171  static_assert(std::is_base_of_v<Remark, RemarkT>,172                "RemarkT must derive from Remark");173  return InFlightRemark(*this,174                        std::make_unique<RemarkT>(std::forward<Args>(args)...));175}176 177template <typename RemarkT>178InFlightRemark179RemarkEngine::emitIfEnabled(Location loc, RemarkOpts opts,180                            bool (RemarkEngine::*isEnabled)(StringRef) const) {181  return (this->*isEnabled)(opts.categoryName) ? makeRemark<RemarkT>(loc, opts)182                                               : InFlightRemark{};183}184 185bool RemarkEngine::isMissedOptRemarkEnabled(StringRef categoryName) const {186  return missFilter && missFilter->match(categoryName);187}188 189bool RemarkEngine::isPassedOptRemarkEnabled(StringRef categoryName) const {190  return passedFilter && passedFilter->match(categoryName);191}192 193bool RemarkEngine::isAnalysisOptRemarkEnabled(StringRef categoryName) const {194  return analysisFilter && analysisFilter->match(categoryName);195}196 197bool RemarkEngine::isFailedOptRemarkEnabled(StringRef categoryName) const {198  return failedFilter && failedFilter->match(categoryName);199}200 201InFlightRemark RemarkEngine::emitOptimizationRemark(Location loc,202                                                    RemarkOpts opts) {203  return emitIfEnabled<OptRemarkPass>(loc, opts,204                                      &RemarkEngine::isPassedOptRemarkEnabled);205}206 207InFlightRemark RemarkEngine::emitOptimizationRemarkMiss(Location loc,208                                                        RemarkOpts opts) {209  return emitIfEnabled<OptRemarkMissed>(210      loc, opts, &RemarkEngine::isMissedOptRemarkEnabled);211}212 213InFlightRemark RemarkEngine::emitOptimizationRemarkFailure(Location loc,214                                                           RemarkOpts opts) {215  return emitIfEnabled<OptRemarkFailure>(216      loc, opts, &RemarkEngine::isFailedOptRemarkEnabled);217}218 219InFlightRemark RemarkEngine::emitOptimizationRemarkAnalysis(Location loc,220                                                            RemarkOpts opts) {221  return emitIfEnabled<OptRemarkAnalysis>(222      loc, opts, &RemarkEngine::isAnalysisOptRemarkEnabled);223}224 225//===----------------------------------------------------------------------===//226// RemarkEngine227//===----------------------------------------------------------------------===//228 229void RemarkEngine::reportImpl(const Remark &remark) {230  // Stream the remark231  if (remarkStreamer) {232    remarkStreamer->streamOptimizationRemark(remark);233  }234 235  // Print using MLIR's diagnostic236  if (printAsEmitRemarks)237    emitRemark(remark.getLocation(), remark.getMsg());238}239 240void RemarkEngine::report(const Remark &&remark) {241  if (remarkEmittingPolicy)242    remarkEmittingPolicy->reportRemark(remark);243}244 245RemarkEngine::~RemarkEngine() {246  if (remarkEmittingPolicy)247    remarkEmittingPolicy->finalize();248 249  if (remarkStreamer)250    remarkStreamer->finalize();251}252 253llvm::LogicalResult RemarkEngine::initialize(254    std::unique_ptr<MLIRRemarkStreamerBase> streamer,255    std::unique_ptr<RemarkEmittingPolicyBase> remarkEmittingPolicy,256    std::string *errMsg) {257 258  remarkStreamer = std::move(streamer);259 260  auto reportFunc =261      std::bind(&RemarkEngine::reportImpl, this, std::placeholders::_1);262  remarkEmittingPolicy->initialize(ReportFn(std::move(reportFunc)));263 264  this->remarkEmittingPolicy = std::move(remarkEmittingPolicy);265  return success();266}267 268/// Returns true if filter is already anchored like ^...$269static bool isAnchored(llvm::StringRef s) {270  s = s.trim();271  return s.starts_with("^") && s.ends_with("$"); // note: startswith/endswith272}273 274/// Anchor the entire pattern so it matches the whole string.275static std::string anchorWhole(llvm::StringRef filter) {276  if (isAnchored(filter))277    return filter.str();278  return (llvm::Twine("^(") + filter + ")$").str();279}280 281/// Build a combined filter from cats.all and a category-specific pattern.282/// If neither is present, return std::nullopt. Otherwise "(all|specific)"283/// and anchor once. Also validate before returning.284static std::optional<llvm::Regex>285buildFilter(const mlir::remark::RemarkCategories &cats,286            const std::optional<std::string> &specific) {287  llvm::SmallVector<llvm::StringRef, 2> parts;288  if (cats.all && !cats.all->empty())289    parts.emplace_back(*cats.all);290  if (specific && !specific->empty())291    parts.emplace_back(*specific);292 293  if (parts.empty())294    return std::nullopt;295 296  std::string joined = llvm::join(parts, "|");297  std::string anchored = anchorWhole(joined);298 299  llvm::Regex rx(anchored);300  std::string err;301  if (!rx.isValid(err))302    return std::nullopt;303 304  return std::make_optional<llvm::Regex>(std::move(rx));305}306 307RemarkEngine::RemarkEngine(bool printAsEmitRemarks,308                           const RemarkCategories &cats)309    : printAsEmitRemarks(printAsEmitRemarks) {310  if (cats.passed)311    passedFilter = buildFilter(cats, cats.passed);312  if (cats.missed)313    missFilter = buildFilter(cats, cats.missed);314  if (cats.analysis)315    analysisFilter = buildFilter(cats, cats.analysis);316  if (cats.failed)317    failedFilter = buildFilter(cats, cats.failed);318}319 320llvm::LogicalResult mlir::remark::enableOptimizationRemarks(321    MLIRContext &ctx, std::unique_ptr<detail::MLIRRemarkStreamerBase> streamer,322    std::unique_ptr<detail::RemarkEmittingPolicyBase> remarkEmittingPolicy,323    const RemarkCategories &cats, bool printAsEmitRemarks) {324  auto engine =325      std::make_unique<detail::RemarkEngine>(printAsEmitRemarks, cats);326 327  std::string errMsg;328  if (failed(engine->initialize(std::move(streamer),329                                std::move(remarkEmittingPolicy), &errMsg))) {330    llvm::report_fatal_error(331        llvm::Twine("Failed to initialize remark engine. Error: ") + errMsg);332  }333  ctx.setRemarkEngine(std::move(engine));334 335  return success();336}337 338//===----------------------------------------------------------------------===//339// Remark emitting policies340//===----------------------------------------------------------------------===//341 342namespace mlir::remark {343RemarkEmittingPolicyAll::RemarkEmittingPolicyAll() = default;344RemarkEmittingPolicyFinal::RemarkEmittingPolicyFinal() = default;345} // namespace mlir::remark346